client.py 11 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301
  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(iteration, result, raw_response=response)
  103. return result
  104. async def achat(
  105. self,
  106. messages: list[Message],
  107. tools: list[ToolDefinition] | None = None,
  108. temperature: float | None = None,
  109. *,
  110. iteration: int = 0,
  111. ) -> Message:
  112. """Async chat completion."""
  113. temp = temperature if temperature is not None else self.settings.agent_temperature
  114. if self.logger:
  115. self.logger.log_llm_input(iteration, self.model, messages, tools, temp)
  116. kwargs: dict[str, Any] = {
  117. "model": self.model,
  118. "messages": [m.to_api_dict() for m in messages],
  119. "tools": [t.to_api_dict() for t in tools] if tools else None,
  120. "temperature": temp,
  121. "timeout": getattr(self.settings, "openrouter_timeout_seconds", 120.0),
  122. }
  123. extra_body = self._extra_body()
  124. if extra_body:
  125. kwargs["extra_body"] = extra_body
  126. response = None
  127. for attempt in range(_MAX_MALFORMED_RETRIES + 1):
  128. response = await self._async_client.chat.completions.create(**kwargs)
  129. if not self._is_malformed_function_call(response):
  130. break
  131. if attempt >= _MAX_MALFORMED_RETRIES:
  132. raise MalformedFunctionCallError(
  133. "LLM provider returned MALFORMED_FUNCTION_CALL "
  134. f"after {_MAX_MALFORMED_RETRIES + 1} attempts"
  135. )
  136. logger.warning(
  137. "LLM provider returned MALFORMED_FUNCTION_CALL; retrying (%d/%d)",
  138. attempt + 1,
  139. _MAX_MALFORMED_RETRIES,
  140. )
  141. kwargs["temperature"] = 0
  142. kwargs["parallel_tool_calls"] = False
  143. kwargs["messages"] = [
  144. *[m.to_api_dict() for m in messages],
  145. {"role": "user", "content": _MALFORMED_RETRY_INSTRUCTION},
  146. ]
  147. assert response is not None
  148. raw_message = response.choices[0].message
  149. result = self._parse_response(raw_message)
  150. if self.logger:
  151. self.logger.log_llm_output(iteration, result, raw_response=response)
  152. return result
  153. def stream(
  154. self,
  155. messages: list[Message],
  156. tools: list[ToolDefinition] | None = None,
  157. temperature: float | None = None,
  158. *,
  159. iteration: int = 0,
  160. ) -> Iterator[str]:
  161. """Stream text content chunks from the model."""
  162. temp = temperature if temperature is not None else self.settings.agent_temperature
  163. if self.logger:
  164. self.logger.log_llm_input(iteration, self.model, messages, tools, temp)
  165. kwargs: dict[str, Any] = {
  166. "model": self.model,
  167. "messages": [m.to_api_dict() for m in messages],
  168. "tools": [t.to_api_dict() for t in tools] if tools else None,
  169. "temperature": temp,
  170. "stream": True,
  171. "timeout": getattr(self.settings, "openrouter_timeout_seconds", 120.0),
  172. }
  173. extra_body = self._extra_body()
  174. if extra_body:
  175. kwargs["extra_body"] = extra_body
  176. stream = self._client.chat.completions.create(**kwargs)
  177. chunks: list[str] = []
  178. for chunk in stream:
  179. delta = chunk.choices[0].delta
  180. if delta.content:
  181. chunks.append(delta.content)
  182. yield delta.content
  183. if self.logger:
  184. full_content = "".join(chunks)
  185. self.logger.log_llm_output(
  186. iteration,
  187. Message(role=Role.ASSISTANT, content=full_content),
  188. )
  189. async def astream(
  190. self,
  191. messages: list[Message],
  192. tools: list[ToolDefinition] | None = None,
  193. temperature: float | None = None,
  194. *,
  195. iteration: int = 0,
  196. ) -> AsyncIterator[str]:
  197. """Async stream text content chunks."""
  198. temp = temperature if temperature is not None else self.settings.agent_temperature
  199. if self.logger:
  200. self.logger.log_llm_input(iteration, self.model, messages, tools, temp)
  201. kwargs: dict[str, Any] = {
  202. "model": self.model,
  203. "messages": [m.to_api_dict() for m in messages],
  204. "tools": [t.to_api_dict() for t in tools] if tools else None,
  205. "temperature": temp,
  206. "stream": True,
  207. "timeout": getattr(self.settings, "openrouter_timeout_seconds", 120.0),
  208. }
  209. extra_body = self._extra_body()
  210. if extra_body:
  211. kwargs["extra_body"] = extra_body
  212. stream = await self._async_client.chat.completions.create(**kwargs)
  213. chunks: list[str] = []
  214. async for chunk in stream:
  215. delta = chunk.choices[0].delta
  216. if delta.content:
  217. chunks.append(delta.content)
  218. yield delta.content
  219. if self.logger:
  220. full_content = "".join(chunks)
  221. self.logger.log_llm_output(
  222. iteration,
  223. Message(role=Role.ASSISTANT, content=full_content),
  224. )
  225. def _parse_response(self, choice_message: Any) -> Message:
  226. tool_calls = None
  227. if choice_message.tool_calls:
  228. tool_calls = [
  229. ToolCall(
  230. id=tc.id,
  231. name=tc.function.name,
  232. arguments=tc.function.arguments,
  233. )
  234. for tc in choice_message.tool_calls
  235. ]
  236. reasoning = getattr(choice_message, "reasoning", None)
  237. if reasoning is not None and not isinstance(reasoning, str):
  238. reasoning = str(reasoning)
  239. return Message(
  240. role=Role.ASSISTANT,
  241. content=choice_message.content,
  242. tool_calls=tool_calls,
  243. reasoning=reasoning,
  244. )
  245. @staticmethod
  246. def _is_malformed_function_call(response: Any) -> bool:
  247. """Detect OpenRouter provider errors that otherwise look like empty answers."""
  248. choices = getattr(response, "choices", None)
  249. if not choices:
  250. return False
  251. choice = choices[0]
  252. native_reason = getattr(choice, "native_finish_reason", None)
  253. if not native_reason:
  254. extra = getattr(choice, "model_extra", None)
  255. if isinstance(extra, dict):
  256. native_reason = extra.get("native_finish_reason")
  257. return str(native_reason or "").upper() == _MALFORMED_FUNCTION_CALL
  258. def set_model(self, model: str) -> None:
  259. """Switch to a different model at runtime."""
  260. self.model = model