champ-ed / agent /agent.py
MalikS-343
squash
cbfe36d
Raw
History Blame Contribute Delete
6.91 kB
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"
@dataclass
class ToolCallRecord:
id: str
function_name: str
arguments: dict
result: str
@dataclass
class AgentResponse:
content: str
tool_calls: list[ToolCallRecord] = field(default_factory=list)
n_tokens: int = 0
@property
def activated_skills(self) -> list[ToolCallRecord]:
return [tc for tc in self.tool_calls if tc.function_name == "activate_skill"]
@property
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,
)