- Added different AIs for ASR and for structure builder
- Forcing models to use JSON from now on
This commit is contained in:
2026-09-23 00:22:09 +03:00
parent 3ae4d82f14
commit 642542ad10
4 changed files with 34 additions and 17 deletions

View File

@@ -1,5 +1,6 @@
from dataclasses import dataclass from dataclasses import dataclass
from typing import Literal from typing import Literal
from pydantic import BaseModel
import httpx2 import httpx2
from openai import OpenAI from openai import OpenAI
@@ -35,7 +36,7 @@ class Agent:
"content": m.content "content": m.content
} }
) )
response = self._client.chat.completions.create( response = self._client.chat.completions.parse(
model=self._model, model=self._model,
messages=messages_raw, messages=messages_raw,
**kwargs **kwargs

View File

@@ -52,7 +52,7 @@ class AsrEventizer:
retries_left = 5 retries_left = 5
while retries_left > 0: while retries_left > 0:
retries_left -= 1 retries_left -= 1
response = self._agent.completion(messages=messages) response = self._agent.completion(messages=messages, response_format=_EventizeResult)
if debug_dir: if debug_dir:
with open(f"{debug_dir}/{retries_left}-retries-left.txt", "w") as f: with open(f"{debug_dir}/{retries_left}-retries-left.txt", "w") as f:
f.write(response) f.write(response)

37
main.py
View File

@@ -50,6 +50,7 @@ def setup_arguments() -> argparse.Namespace:
- argparse namespace - argparse namespace
""" """
voice_models = Asr.get_models_list() voice_models = Asr.get_models_list()
ai_api_key = os.environ.get("AI_API_KEY", default=...)
parser = argparse.ArgumentParser( parser = argparse.ArgumentParser(
prog="sumka", prog="sumka",
@@ -66,18 +67,34 @@ def setup_arguments() -> argparse.Namespace:
default="ru" default="ru"
) )
parser.add_argument( parser.add_argument(
"--ai-model", "--asr-ai-model",
type=str, type=str,
default="google/gemini-3.1-flash-lite" default="google/gemini-3.1-flash-lite"
) )
parser.add_argument( parser.add_argument(
"--ai-base-url", "--asr-ai-base-url",
type=str, type=str,
default="https://api.proxyapi.ru/v1" default="https://api.proxyapi.ru/v1"
) )
parser.add_argument( parser.add_argument(
"--ai-api-key", "--asr-ai-api-key",
type=str type=str,
default=ai_api_key
)
parser.add_argument(
"--structure-ai-model",
type=str,
default="google/gemini-3.5-flash"
)
parser.add_argument(
"--structure-ai-base-url",
type=str,
default="https://api.proxyapi.ru/v1"
)
parser.add_argument(
"--structure-ai-api-key",
type=str,
default=ai_api_key
) )
parser.add_argument("-v", action='store_true') parser.add_argument("-v", action='store_true')
return parser.parse_args() return parser.parse_args()
@@ -160,9 +177,9 @@ def on_asr_events(current_step: Step, input_data: dict | None) -> tuple[Step | N
return (None, None) return (None, None)
logging.info("Creating the agent") logging.info("Creating the agent")
agent = Agent( agent = Agent(
model=ARGS.ai_model, model=ARGS.asr_ai_model,
base_url=ARGS.ai_base_url, base_url=ARGS.asr_ai_base_url,
api_key=ARGS.ai_api_key api_key=ARGS.asr_ai_api_key
) )
logging.info("Creating the windowizer") logging.info("Creating the windowizer")
windowizer = Windowizer() windowizer = Windowizer()
@@ -208,9 +225,9 @@ def on_structure_builder(current_step: Step, input_data: dict | None) -> tuple[S
return (None, None) return (None, None)
logging.info("Creating the agent") logging.info("Creating the agent")
agent = Agent( agent = Agent(
model=ARGS.ai_model, model=ARGS.structure_ai_model,
base_url=ARGS.ai_base_url, base_url=ARGS.structure_ai_base_url,
api_key=ARGS.ai_api_key api_key=ARGS.structure_ai_api_key
) )
logging.info("Creating the windowizer") logging.info("Creating the windowizer")
windowizer = Windowizer() windowizer = Windowizer()

View File

@@ -79,16 +79,15 @@ class ImageElement(_StrictModel):
event_id: int event_id: int
StructureElement = Annotated[ StructureElement = (
HeadingElement HeadingElement
| ParagraphElement | ParagraphElement
| UnorderedListElement | UnorderedListElement
| OrderedListElement | OrderedListElement
| DefinitionElement | DefinitionElement
| ImportantElement | ImportantElement
| ImageElement, | ImageElement
Field(discriminator="type"), )
]
class Structure(_StrictModel): class Structure(_StrictModel):
@@ -204,7 +203,7 @@ class StructureBuilder:
retries_left = self.MAX_RETRIES retries_left = self.MAX_RETRIES
while retries_left > 0: while retries_left > 0:
retries_left -= 1 retries_left -= 1
response = self._agent.completion(messages=messages) response = self._agent.completion(messages=messages, response_format=_BuildResult)
if debug_dir: if debug_dir:
with open( with open(
f"{debug_dir}/{retries_left}-retries-left.txt", f"{debug_dir}/{retries_left}-retries-left.txt",