| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301 |
- 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
|