File size: 5,816 Bytes
c63bc31
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
"""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:"))