from __future__ import annotations import logging from collections.abc import AsyncIterator, Iterator from typing import TYPE_CHECKING, Any from openai import AsyncOpenAI, OpenAI from supply_agent.config import Settings from supply_agent.types import Message, Role, ToolCall, ToolDefinition if TYPE_CHECKING: from supply_agent.logging.logger import AgentLogger logger = logging.getLogger(__name__) _MALFORMED_FUNCTION_CALL = "MALFORMED_FUNCTION_CALL" _MAX_MALFORMED_RETRIES = 3 _MALFORMED_RETRY_INSTRUCTION = ( "上一次工具调用不是有效 JSON。请重新决定下一步,只调用一个最必要的工具;" "严格使用工具 schema,省略未定义字段,并确保 arguments 是完整 JSON。" ) class MalformedFunctionCallError(RuntimeError): """Provider repeatedly failed to produce a valid tool call.""" class LLMClient: """OpenRouter LLM client using the OpenAI-compatible API.""" def __init__( self, settings: Settings, logger: AgentLogger | None = None, *, reasoning_effort: str | None = None, ) -> None: self.settings = settings self.model = settings.openrouter_model self.logger = logger # OpenRouter reasoning effort: low / medium / high / max / xhigh. # None or "" disables the parameter. self.reasoning_effort = reasoning_effort extra_headers: dict[str, str] = {} if settings.openrouter_site_url: extra_headers["HTTP-Referer"] = settings.openrouter_site_url if settings.openrouter_site_name: extra_headers["X-Title"] = settings.openrouter_site_name self._client = OpenAI( api_key=settings.openrouter_api_key, base_url=settings.openrouter_base_url, default_headers=extra_headers or None, ) self._async_client = AsyncOpenAI( api_key=settings.openrouter_api_key, base_url=settings.openrouter_base_url, default_headers=extra_headers or None, ) def _extra_body(self) -> dict[str, Any] | None: """OpenRouter-only params (e.g. Claude reasoning effort).""" effort = (self.reasoning_effort or "").strip() if not effort: return None return {"reasoning": {"effort": effort}} def chat( self, messages: list[Message], tools: list[ToolDefinition] | None = None, temperature: float | None = None, *, iteration: int = 0, ) -> Message: """Send a chat completion request and return the assistant message.""" temp = temperature if temperature is not None else self.settings.agent_temperature if self.logger: self.logger.log_llm_input(iteration, self.model, messages, tools, temp) kwargs: dict[str, Any] = { "model": self.model, "messages": [m.to_api_dict() for m in messages], "tools": [t.to_api_dict() for t in tools] if tools else None, "temperature": temp, "timeout": getattr(self.settings, "openrouter_timeout_seconds", 120.0), } extra_body = self._extra_body() if extra_body: kwargs["extra_body"] = extra_body response = None for attempt in range(_MAX_MALFORMED_RETRIES + 1): response = self._client.chat.completions.create(**kwargs) if not self._is_malformed_function_call(response): break if attempt >= _MAX_MALFORMED_RETRIES: raise MalformedFunctionCallError( "LLM provider returned MALFORMED_FUNCTION_CALL " f"after {_MAX_MALFORMED_RETRIES + 1} attempts" ) logger.warning( "LLM provider returned MALFORMED_FUNCTION_CALL; retrying (%d/%d)", attempt + 1, _MAX_MALFORMED_RETRIES, ) kwargs["temperature"] = 0 kwargs["parallel_tool_calls"] = False kwargs["messages"] = [ *[m.to_api_dict() for m in messages], {"role": "user", "content": _MALFORMED_RETRY_INSTRUCTION}, ] assert response is not None raw_message = response.choices[0].message result = self._parse_response(raw_message) if self.logger: self.logger.log_llm_output(iteration, result, raw_response=response) return result async def achat( self, messages: list[Message], tools: list[ToolDefinition] | None = None, temperature: float | None = None, *, iteration: int = 0, ) -> Message: """Async chat completion.""" temp = temperature if temperature is not None else self.settings.agent_temperature if self.logger: self.logger.log_llm_input(iteration, self.model, messages, tools, temp) kwargs: dict[str, Any] = { "model": self.model, "messages": [m.to_api_dict() for m in messages], "tools": [t.to_api_dict() for t in tools] if tools else None, "temperature": temp, "timeout": getattr(self.settings, "openrouter_timeout_seconds", 120.0), } extra_body = self._extra_body() if extra_body: kwargs["extra_body"] = extra_body response = None for attempt in range(_MAX_MALFORMED_RETRIES + 1): response = await self._async_client.chat.completions.create(**kwargs) if not self._is_malformed_function_call(response): break if attempt >= _MAX_MALFORMED_RETRIES: raise MalformedFunctionCallError( "LLM provider returned MALFORMED_FUNCTION_CALL " f"after {_MAX_MALFORMED_RETRIES + 1} attempts" ) logger.warning( "LLM provider returned MALFORMED_FUNCTION_CALL; retrying (%d/%d)", attempt + 1, _MAX_MALFORMED_RETRIES, ) kwargs["temperature"] = 0 kwargs["parallel_tool_calls"] = False kwargs["messages"] = [ *[m.to_api_dict() for m in messages], {"role": "user", "content": _MALFORMED_RETRY_INSTRUCTION}, ] assert response is not None raw_message = response.choices[0].message result = self._parse_response(raw_message) if self.logger: self.logger.log_llm_output(iteration, result, raw_response=response) return result def stream( self, messages: list[Message], tools: list[ToolDefinition] | None = None, temperature: float | None = None, *, iteration: int = 0, ) -> Iterator[str]: """Stream text content chunks from the model.""" temp = temperature if temperature is not None else self.settings.agent_temperature if self.logger: self.logger.log_llm_input(iteration, self.model, messages, tools, temp) kwargs: dict[str, Any] = { "model": self.model, "messages": [m.to_api_dict() for m in messages], "tools": [t.to_api_dict() for t in tools] if tools else None, "temperature": temp, "stream": True, "timeout": getattr(self.settings, "openrouter_timeout_seconds", 120.0), } extra_body = self._extra_body() if extra_body: kwargs["extra_body"] = extra_body stream = self._client.chat.completions.create(**kwargs) chunks: list[str] = [] for chunk in stream: delta = chunk.choices[0].delta if delta.content: chunks.append(delta.content) yield delta.content if self.logger: full_content = "".join(chunks) self.logger.log_llm_output( iteration, Message(role=Role.ASSISTANT, content=full_content), ) async def astream( self, messages: list[Message], tools: list[ToolDefinition] | None = None, temperature: float | None = None, *, iteration: int = 0, ) -> AsyncIterator[str]: """Async stream text content chunks.""" temp = temperature if temperature is not None else self.settings.agent_temperature if self.logger: self.logger.log_llm_input(iteration, self.model, messages, tools, temp) kwargs: dict[str, Any] = { "model": self.model, "messages": [m.to_api_dict() for m in messages], "tools": [t.to_api_dict() for t in tools] if tools else None, "temperature": temp, "stream": True, "timeout": getattr(self.settings, "openrouter_timeout_seconds", 120.0), } extra_body = self._extra_body() if extra_body: kwargs["extra_body"] = extra_body stream = await self._async_client.chat.completions.create(**kwargs) chunks: list[str] = [] async for chunk in stream: delta = chunk.choices[0].delta if delta.content: chunks.append(delta.content) yield delta.content if self.logger: full_content = "".join(chunks) self.logger.log_llm_output( iteration, Message(role=Role.ASSISTANT, content=full_content), ) def _parse_response(self, choice_message: Any) -> Message: tool_calls = None if choice_message.tool_calls: tool_calls = [ ToolCall( id=tc.id, name=tc.function.name, arguments=tc.function.arguments, ) for tc in choice_message.tool_calls ] reasoning = getattr(choice_message, "reasoning", None) if reasoning is not None and not isinstance(reasoning, str): reasoning = str(reasoning) return Message( role=Role.ASSISTANT, content=choice_message.content, tool_calls=tool_calls, reasoning=reasoning, ) @staticmethod def _is_malformed_function_call(response: Any) -> bool: """Detect OpenRouter provider errors that otherwise look like empty answers.""" choices = getattr(response, "choices", None) if not choices: return False choice = choices[0] native_reason = getattr(choice, "native_finish_reason", None) if not native_reason: extra = getattr(choice, "model_extra", None) if isinstance(extra, dict): native_reason = extra.get("native_finish_reason") return str(native_reason or "").upper() == _MALFORMED_FUNCTION_CALL def set_model(self, model: str) -> None: """Switch to a different model at runtime.""" self.model = model