Spaces:
Running on CPU Upgrade
Running on CPU Upgrade
| from dataclasses import dataclass, field | |
| import logging | |
| from agent.agent_exceptions import InferenceError, MaxTurnsExceeded | |
| from agent.conversation_history import ConversationHistory | |
| from agent.documents import build_documents_block | |
| from agent.skill_manager import SkillsManager | |
| from agent.prompts import SKILLS_SYSTEM_PROMPT | |
| from helpers.timing import timed_block | |
| from providers.exceptions import ProviderError | |
| from providers.protocol import ChatProvider, Message, ReasoningEffort, ToolCall | |
| logger = logging.getLogger(__name__) | |
| DEFAULT_AGENT_MODEL = "openai/gpt-oss-20b" | |
| DEFAULT_AGENT_MAX_TURNS = 10 | |
| DEFAULT_REASONING_EFFORT: ReasoningEffort = "medium" | |
| class ToolCallRecord: | |
| id: str | |
| function_name: str | |
| arguments: dict | |
| result: str | |
| class AgentResponse: | |
| content: str | |
| tool_calls: list[ToolCallRecord] = field(default_factory=list) | |
| n_tokens: int = 0 | |
| def activated_skills(self) -> list[ToolCallRecord]: | |
| return [tc for tc in self.tool_calls if tc.function_name == "activate_skill"] | |
| def executed_functions(self) -> list[ToolCallRecord]: | |
| return [tc for tc in self.tool_calls if tc.function_name == "execute_function"] | |
| def to_judge_context(self) -> dict: | |
| return { | |
| "agent_response": self.content, | |
| "tool_calls": [ | |
| { | |
| "function_name": tc.function_name, | |
| "arguments": tc.arguments, | |
| "result": tc.result, | |
| } | |
| for tc in self.tool_calls | |
| ], | |
| } | |
| class Agent: | |
| def __init__( | |
| self, | |
| skills: SkillsManager, | |
| provider: ChatProvider, | |
| model_id: str = DEFAULT_AGENT_MODEL, | |
| system_prompt: str = SKILLS_SYSTEM_PROMPT, | |
| reasoning_effort: ReasoningEffort | None = DEFAULT_REASONING_EFFORT, | |
| max_turns: int = DEFAULT_AGENT_MAX_TURNS, | |
| ) -> None: | |
| self.skills = skills | |
| self.provider = provider | |
| self.model_id = model_id | |
| self.system_prompt = system_prompt | |
| self.reasoning_effort = reasoning_effort | |
| self.max_turns = max_turns | |
| def chat( | |
| self, | |
| query: str, | |
| conversation: ConversationHistory, | |
| *, | |
| system: bool = False, | |
| documents: dict[str, str] | None = None, | |
| ) -> AgentResponse: | |
| if system: | |
| conversation.record_system(query) | |
| else: | |
| conversation.record_user(query) | |
| tool_call_records: list[ToolCallRecord] = [] | |
| n_tokens = 0 | |
| for _ in range(self.max_turns): | |
| with timed_block("agent.llm_call"): | |
| try: | |
| completion = self.provider.chat( | |
| messages=self._build_messages(conversation, documents), | |
| model_id=self.model_id, | |
| tools=self.skills.tools_for(conversation), | |
| reasoning_effort=self.reasoning_effort, | |
| ) | |
| except ProviderError as e: | |
| raise InferenceError(str(e)) from e | |
| n_tokens += completion.usage.total_tokens | |
| msg = completion.message | |
| reasoning = msg.reasoning | |
| if msg.tool_calls: | |
| for tc in msg.tool_calls: | |
| try: | |
| record, should_return = self._dispatch_tool_call( | |
| tc, conversation, documents | |
| ) | |
| except Exception as exc: | |
| conversation.record_tool_exchange( | |
| tool_call_id=tc.id, | |
| function_name=tc.name, | |
| arguments=tc.arguments, | |
| result=f"[DISPATCH ERROR] {type(exc).__name__}: {exc}", | |
| reasoning=reasoning, | |
| ) | |
| raise | |
| conversation.record_tool_exchange( | |
| tool_call_id=record.id, | |
| function_name=record.function_name, | |
| arguments=record.arguments, | |
| result=record.result, | |
| reasoning=reasoning, | |
| ) | |
| tool_call_records.append(record) | |
| if should_return: | |
| conversation.record_assistant(record.result, reasoning=None) | |
| return AgentResponse( | |
| content=record.result, | |
| tool_calls=tool_call_records, | |
| n_tokens=n_tokens, | |
| ) | |
| continue | |
| if msg.content is None: | |
| raise ValueError("Provider returned neither tool_calls nor content") | |
| conversation.record_assistant(msg.content, reasoning=reasoning) | |
| return AgentResponse( | |
| content=msg.content, tool_calls=tool_call_records, n_tokens=n_tokens | |
| ) | |
| raise MaxTurnsExceeded(f"Exceeded {self.max_turns} turns without a final reply") | |
| def _build_messages( | |
| self, | |
| conversation: ConversationHistory, | |
| documents: dict[str, str] | None = None, | |
| ) -> list[Message]: | |
| system_content = self._build_system_prompt() | |
| docs_block = build_documents_block(documents) | |
| if docs_block: | |
| system_content = f"{system_content}\n\n{docs_block}" | |
| return [ | |
| Message(role="system", content=system_content), | |
| *conversation.to_messages(), | |
| ] | |
| def _build_system_prompt(self) -> str: | |
| return self.system_prompt.format( | |
| skill_list=self.skills.to_system_prompt_format() | |
| ) | |
| def _dispatch_tool_call( | |
| self, | |
| tc: ToolCall, | |
| conversation: ConversationHistory, | |
| documents: dict[str, str] | None = None, | |
| ) -> tuple[ToolCallRecord, bool]: | |
| should_return = False | |
| try: | |
| if tc.name == "activate_skill": | |
| instructions = self.skills.activate(**tc.arguments) | |
| result = f"Instructions: {instructions}" | |
| elif tc.name == "execute_function": | |
| result, should_return = self.skills.execute( | |
| **tc.arguments, conversation=conversation, documents=documents | |
| ) | |
| else: | |
| result = f"Error: Unknown function: {tc.name}" | |
| except TypeError as e: | |
| result = f"An unexpected keyword argument was passed to the function you were trying to call: {e}" | |
| logger.warning( | |
| "Model generated an unexpected keyword argument for a tool call." | |
| ) | |
| return ( | |
| ToolCallRecord( | |
| id=tc.id, function_name=tc.name, arguments=tc.arguments, result=result | |
| ), | |
| should_return, | |
| ) | |