client.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305
  1. from __future__ import annotations
  2. import logging
  3. from collections.abc import AsyncIterator, Iterator
  4. from typing import TYPE_CHECKING, Any
  5. from openai import AsyncOpenAI, OpenAI
  6. from supply_agent.config import Settings
  7. from supply_agent.types import Message, Role, ToolCall, ToolDefinition
  8. if TYPE_CHECKING:
  9. from supply_agent.logging.logger import AgentLogger
  10. logger = logging.getLogger(__name__)
  11. _MALFORMED_FUNCTION_CALL = "MALFORMED_FUNCTION_CALL"
  12. _MAX_MALFORMED_RETRIES = 3
  13. _MALFORMED_RETRY_INSTRUCTION = (
  14. "上一次工具调用不是有效 JSON。请重新决定下一步,只调用一个最必要的工具;"
  15. "严格使用工具 schema,省略未定义字段,并确保 arguments 是完整 JSON。"
  16. )
  17. class MalformedFunctionCallError(RuntimeError):
  18. """Provider repeatedly failed to produce a valid tool call."""
  19. class LLMClient:
  20. """OpenRouter LLM client using the OpenAI-compatible API."""
  21. def __init__(
  22. self,
  23. settings: Settings,
  24. logger: AgentLogger | None = None,
  25. *,
  26. reasoning_effort: str | None = None,
  27. ) -> None:
  28. self.settings = settings
  29. self.model = settings.openrouter_model
  30. self.logger = logger
  31. # OpenRouter reasoning effort: low / medium / high / max / xhigh.
  32. # None or "" disables the parameter.
  33. self.reasoning_effort = reasoning_effort
  34. extra_headers: dict[str, str] = {}
  35. if settings.openrouter_site_url:
  36. extra_headers["HTTP-Referer"] = settings.openrouter_site_url
  37. if settings.openrouter_site_name:
  38. extra_headers["X-Title"] = settings.openrouter_site_name
  39. self._client = OpenAI(
  40. api_key=settings.openrouter_api_key,
  41. base_url=settings.openrouter_base_url,
  42. default_headers=extra_headers or None,
  43. )
  44. self._async_client = AsyncOpenAI(
  45. api_key=settings.openrouter_api_key,
  46. base_url=settings.openrouter_base_url,
  47. default_headers=extra_headers or None,
  48. )
  49. def _extra_body(self) -> dict[str, Any] | None:
  50. """OpenRouter-only params (e.g. Claude reasoning effort)."""
  51. effort = (self.reasoning_effort or "").strip()
  52. if not effort:
  53. return None
  54. return {"reasoning": {"effort": effort}}
  55. def chat(
  56. self,
  57. messages: list[Message],
  58. tools: list[ToolDefinition] | None = None,
  59. temperature: float | None = None,
  60. *,
  61. iteration: int = 0,
  62. ) -> Message:
  63. """Send a chat completion request and return the assistant message."""
  64. temp = temperature if temperature is not None else self.settings.agent_temperature
  65. if self.logger:
  66. self.logger.log_llm_input(iteration, self.model, messages, tools, temp)
  67. kwargs: dict[str, Any] = {
  68. "model": self.model,
  69. "messages": [m.to_api_dict() for m in messages],
  70. "tools": [t.to_api_dict() for t in tools] if tools else None,
  71. "temperature": temp,
  72. "timeout": getattr(self.settings, "openrouter_timeout_seconds", 120.0),
  73. }
  74. extra_body = self._extra_body()
  75. if extra_body:
  76. kwargs["extra_body"] = extra_body
  77. response = None
  78. for attempt in range(_MAX_MALFORMED_RETRIES + 1):
  79. response = self._client.chat.completions.create(**kwargs)
  80. if not self._is_malformed_function_call(response):
  81. break
  82. if attempt >= _MAX_MALFORMED_RETRIES:
  83. raise MalformedFunctionCallError(
  84. "LLM provider returned MALFORMED_FUNCTION_CALL "
  85. f"after {_MAX_MALFORMED_RETRIES + 1} attempts"
  86. )
  87. logger.warning(
  88. "LLM provider returned MALFORMED_FUNCTION_CALL; retrying (%d/%d)",
  89. attempt + 1,
  90. _MAX_MALFORMED_RETRIES,
  91. )
  92. kwargs["temperature"] = 0
  93. kwargs["parallel_tool_calls"] = False
  94. kwargs["messages"] = [
  95. *[m.to_api_dict() for m in messages],
  96. {"role": "user", "content": _MALFORMED_RETRY_INSTRUCTION},
  97. ]
  98. assert response is not None
  99. raw_message = response.choices[0].message
  100. result = self._parse_response(raw_message)
  101. if self.logger:
  102. self.logger.log_llm_output(
  103. iteration, result, raw_response=response, model=self.model
  104. )
  105. return result
  106. async def achat(
  107. self,
  108. messages: list[Message],
  109. tools: list[ToolDefinition] | None = None,
  110. temperature: float | None = None,
  111. *,
  112. iteration: int = 0,
  113. ) -> Message:
  114. """Async chat completion."""
  115. temp = temperature if temperature is not None else self.settings.agent_temperature
  116. if self.logger:
  117. self.logger.log_llm_input(iteration, self.model, messages, tools, temp)
  118. kwargs: dict[str, Any] = {
  119. "model": self.model,
  120. "messages": [m.to_api_dict() for m in messages],
  121. "tools": [t.to_api_dict() for t in tools] if tools else None,
  122. "temperature": temp,
  123. "timeout": getattr(self.settings, "openrouter_timeout_seconds", 120.0),
  124. }
  125. extra_body = self._extra_body()
  126. if extra_body:
  127. kwargs["extra_body"] = extra_body
  128. response = None
  129. for attempt in range(_MAX_MALFORMED_RETRIES + 1):
  130. response = await self._async_client.chat.completions.create(**kwargs)
  131. if not self._is_malformed_function_call(response):
  132. break
  133. if attempt >= _MAX_MALFORMED_RETRIES:
  134. raise MalformedFunctionCallError(
  135. "LLM provider returned MALFORMED_FUNCTION_CALL "
  136. f"after {_MAX_MALFORMED_RETRIES + 1} attempts"
  137. )
  138. logger.warning(
  139. "LLM provider returned MALFORMED_FUNCTION_CALL; retrying (%d/%d)",
  140. attempt + 1,
  141. _MAX_MALFORMED_RETRIES,
  142. )
  143. kwargs["temperature"] = 0
  144. kwargs["parallel_tool_calls"] = False
  145. kwargs["messages"] = [
  146. *[m.to_api_dict() for m in messages],
  147. {"role": "user", "content": _MALFORMED_RETRY_INSTRUCTION},
  148. ]
  149. assert response is not None
  150. raw_message = response.choices[0].message
  151. result = self._parse_response(raw_message)
  152. if self.logger:
  153. self.logger.log_llm_output(
  154. iteration, result, raw_response=response, model=self.model
  155. )
  156. return result
  157. def stream(
  158. self,
  159. messages: list[Message],
  160. tools: list[ToolDefinition] | None = None,
  161. temperature: float | None = None,
  162. *,
  163. iteration: int = 0,
  164. ) -> Iterator[str]:
  165. """Stream text content chunks from the model."""
  166. temp = temperature if temperature is not None else self.settings.agent_temperature
  167. if self.logger:
  168. self.logger.log_llm_input(iteration, self.model, messages, tools, temp)
  169. kwargs: dict[str, Any] = {
  170. "model": self.model,
  171. "messages": [m.to_api_dict() for m in messages],
  172. "tools": [t.to_api_dict() for t in tools] if tools else None,
  173. "temperature": temp,
  174. "stream": True,
  175. "timeout": getattr(self.settings, "openrouter_timeout_seconds", 120.0),
  176. }
  177. extra_body = self._extra_body()
  178. if extra_body:
  179. kwargs["extra_body"] = extra_body
  180. stream = self._client.chat.completions.create(**kwargs)
  181. chunks: list[str] = []
  182. for chunk in stream:
  183. delta = chunk.choices[0].delta
  184. if delta.content:
  185. chunks.append(delta.content)
  186. yield delta.content
  187. if self.logger:
  188. full_content = "".join(chunks)
  189. self.logger.log_llm_output(
  190. iteration,
  191. Message(role=Role.ASSISTANT, content=full_content),
  192. )
  193. async def astream(
  194. self,
  195. messages: list[Message],
  196. tools: list[ToolDefinition] | None = None,
  197. temperature: float | None = None,
  198. *,
  199. iteration: int = 0,
  200. ) -> AsyncIterator[str]:
  201. """Async stream text content chunks."""
  202. temp = temperature if temperature is not None else self.settings.agent_temperature
  203. if self.logger:
  204. self.logger.log_llm_input(iteration, self.model, messages, tools, temp)
  205. kwargs: dict[str, Any] = {
  206. "model": self.model,
  207. "messages": [m.to_api_dict() for m in messages],
  208. "tools": [t.to_api_dict() for t in tools] if tools else None,
  209. "temperature": temp,
  210. "stream": True,
  211. "timeout": getattr(self.settings, "openrouter_timeout_seconds", 120.0),
  212. }
  213. extra_body = self._extra_body()
  214. if extra_body:
  215. kwargs["extra_body"] = extra_body
  216. stream = await self._async_client.chat.completions.create(**kwargs)
  217. chunks: list[str] = []
  218. async for chunk in stream:
  219. delta = chunk.choices[0].delta
  220. if delta.content:
  221. chunks.append(delta.content)
  222. yield delta.content
  223. if self.logger:
  224. full_content = "".join(chunks)
  225. self.logger.log_llm_output(
  226. iteration,
  227. Message(role=Role.ASSISTANT, content=full_content),
  228. )
  229. def _parse_response(self, choice_message: Any) -> Message:
  230. tool_calls = None
  231. if choice_message.tool_calls:
  232. tool_calls = [
  233. ToolCall(
  234. id=tc.id,
  235. name=tc.function.name,
  236. arguments=tc.function.arguments,
  237. )
  238. for tc in choice_message.tool_calls
  239. ]
  240. reasoning = getattr(choice_message, "reasoning", None)
  241. if reasoning is not None and not isinstance(reasoning, str):
  242. reasoning = str(reasoning)
  243. return Message(
  244. role=Role.ASSISTANT,
  245. content=choice_message.content,
  246. tool_calls=tool_calls,
  247. reasoning=reasoning,
  248. )
  249. @staticmethod
  250. def _is_malformed_function_call(response: Any) -> bool:
  251. """Detect OpenRouter provider errors that otherwise look like empty answers."""
  252. choices = getattr(response, "choices", None)
  253. if not choices:
  254. return False
  255. choice = choices[0]
  256. native_reason = getattr(choice, "native_finish_reason", None)
  257. if not native_reason:
  258. extra = getattr(choice, "model_extra", None)
  259. if isinstance(extra, dict):
  260. native_reason = extra.get("native_finish_reason")
  261. return str(native_reason or "").upper() == _MALFORMED_FUNCTION_CALL
  262. def set_model(self, model: str) -> None:
  263. """Switch to a different model at runtime."""
  264. self.model = model