asdfasdfqrqwer's picture
feat(cortex-ai): agent engine, tool registry, OpenAI-compatible API, fine-tuning pipeline
c63bc31
Raw History Blame Contribute Delete
5.82 kB
"""The CORTEX agent loop.
Ties together the encoding layer, a model adapter and the tool registry:
messages -> encode_messages -> model -> parse -> tool call? -> execute
-> append tool result -> encode again -> ... -> final answer
The loop is backend-agnostic. With ``MockAdapter`` it runs on a CPU with no
downloads, which is how it is tested in this repository.
"""
from __future__ import annotations
import json
from dataclasses import dataclass, field
from typing import Any
from ..adapters.base import ModelAdapter
from ..config import EngineConfig
from ..tools.registry import ToolError, ToolRegistry
@dataclass
class ToolCallRecord:
"""One tool invocation performed during a turn."""
name: str
arguments: dict[str, Any]
result: str
ok: bool
@dataclass
class AgentTurn:
"""The complete result of one user turn."""
content: str = ""
reasoning: str = ""
tool_calls: list[ToolCallRecord] = field(default_factory=list)
rounds: int = 0
raw_completions: list[str] = field(default_factory=list)
@property
def used_tools(self) -> bool:
return bool(self.tool_calls)
def _encode(messages: list[dict[str, Any]], engine: EngineConfig) -> str:
"""Call the DeepSeek-V4 encoder, with a clear error if it is missing."""
try:
from encoding_dsv4 import encode_messages
except ImportError as exc: # pragma: no cover - environment dependent
raise RuntimeError(
"The encoding module is required. Add the repository's encoding/ folder "
"to PYTHONPATH, e.g. PYTHONPATH=/workspace/project/encoding"
) from exc
return encode_messages(
messages,
thinking_mode=engine.thinking_mode,
reasoning_effort=engine.reasoning_effort,
)
def _parse(text: str, engine: EngineConfig) -> dict[str, Any]:
try:
from encoding_dsv4 import parse_message_from_completion_text
except ImportError as exc: # pragma: no cover
raise RuntimeError("The encoding module is required.") from exc
return parse_message_from_completion_text(text, thinking_mode=engine.thinking_mode)
def _normalize_call(call: dict[str, Any]) -> tuple[str, Any]:
"""Accept both the flat and the OpenAI-nested tool-call shape.
``parse_message_from_completion_text`` returns OpenAI format::
{"type": "function", "function": {"name": ..., "arguments": "<json>"}}
while hand-written or replayed data may use ``{"name": ..., "arguments": ...}``.
"""
if "function" in call and isinstance(call["function"], dict):
fn = call["function"]
return fn.get("name", ""), fn.get("arguments", "{}")
return call.get("name", ""), call.get("arguments", {})
class CortexAgent:
"""Runs the reasoning-and-tool loop for a conversation."""
def __init__(
self,
adapter: ModelAdapter,
tools: ToolRegistry | None = None,
engine: EngineConfig | None = None,
system_prompt: str | None = None,
) -> None:
self.adapter = adapter
self.tools = tools if tools is not None else ToolRegistry()
self.engine = engine or EngineConfig()
self.system_prompt = system_prompt
def _system_message(self) -> dict[str, Any]:
msg: dict[str, Any] = {"role": "system", "content": self.system_prompt or ""}
if len(self.tools):
msg["tools"] = self.tools.to_openai_schemas()
return msg
def run(self, messages: list[dict[str, Any]]) -> AgentTurn:
"""Execute one turn. *messages* are OpenAI-style dicts, without the system message."""
turn = AgentTurn()
convo: list[dict[str, Any]] = [self._system_message(), *messages]
for round_index in range(self.engine.max_tool_rounds):
turn.rounds = round_index + 1
prompt = _encode(convo, self.engine)
result = self.adapter.generate(
prompt,
max_new_tokens=self.engine.max_new_tokens,
temperature=self.engine.temperature,
)
turn.raw_completions.append(result.text)
parsed = _parse(result.text, self.engine)
turn.reasoning = parsed.get("reasoning_content") or turn.reasoning
calls = parsed.get("tool_calls") or []
if not calls:
turn.content = parsed.get("content", "")
return turn
convo.append(
{
"role": "assistant",
"content": parsed.get("content", ""),
"reasoning_content": parsed.get("reasoning_content", ""),
"tool_calls": calls,
}
)
for call in calls:
name, raw_args = _normalize_call(call)
record = self._execute(name, raw_args)
turn.tool_calls.append(record)
convo.append({"role": "tool", "content": record.result})
turn.content = (
f"Limite de {self.engine.max_tool_rounds} tours d'outils atteinte "
"sans réponse finale."
)
return turn
def _execute(self, name: str, raw_args: Any) -> ToolCallRecord:
try:
if isinstance(raw_args, str):
args = json.loads(raw_args) if raw_args.strip() else {}
else:
args = dict(raw_args or {})
except json.JSONDecodeError:
args = {}
return ToolCallRecord(name, {}, f"ERROR: arguments are not valid JSON: {raw_args}", False)
try:
result = self.tools.call(name, args)
except ToolError as exc:
return ToolCallRecord(name, args, f"ERROR: {exc}", False)
return ToolCallRecord(name, args, result, not result.startswith("ERROR:"))