client.py 7.5 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225
  1. from __future__ import annotations
  2. from collections.abc import AsyncIterator, Iterator
  3. from typing import TYPE_CHECKING, Any
  4. from openai import AsyncOpenAI, OpenAI
  5. from supply_agent.config import Settings
  6. from supply_agent.types import Message, Role, ToolCall, ToolDefinition
  7. if TYPE_CHECKING:
  8. from supply_agent.logging.logger import AgentLogger
  9. class LLMClient:
  10. """OpenRouter LLM client using the OpenAI-compatible API."""
  11. def __init__(
  12. self,
  13. settings: Settings,
  14. logger: AgentLogger | None = None,
  15. *,
  16. reasoning_effort: str | None = None,
  17. ) -> None:
  18. self.settings = settings
  19. self.model = settings.openrouter_model
  20. self.logger = logger
  21. # OpenRouter reasoning effort: low / medium / high / max / xhigh.
  22. # None or "" disables the parameter.
  23. self.reasoning_effort = reasoning_effort
  24. extra_headers: dict[str, str] = {}
  25. if settings.openrouter_site_url:
  26. extra_headers["HTTP-Referer"] = settings.openrouter_site_url
  27. if settings.openrouter_site_name:
  28. extra_headers["X-Title"] = settings.openrouter_site_name
  29. self._client = OpenAI(
  30. api_key=settings.openrouter_api_key,
  31. base_url=settings.openrouter_base_url,
  32. default_headers=extra_headers or None,
  33. )
  34. self._async_client = AsyncOpenAI(
  35. api_key=settings.openrouter_api_key,
  36. base_url=settings.openrouter_base_url,
  37. default_headers=extra_headers or None,
  38. )
  39. def _extra_body(self) -> dict[str, Any] | None:
  40. """OpenRouter-only params (e.g. Claude reasoning effort)."""
  41. effort = (self.reasoning_effort or "").strip()
  42. if not effort:
  43. return None
  44. return {"reasoning": {"effort": effort}}
  45. def chat(
  46. self,
  47. messages: list[Message],
  48. tools: list[ToolDefinition] | None = None,
  49. temperature: float | None = None,
  50. *,
  51. iteration: int = 0,
  52. ) -> Message:
  53. """Send a chat completion request and return the assistant message."""
  54. temp = temperature if temperature is not None else self.settings.agent_temperature
  55. if self.logger:
  56. self.logger.log_llm_input(iteration, self.model, messages, tools, temp)
  57. kwargs: dict[str, Any] = {
  58. "model": self.model,
  59. "messages": [m.to_api_dict() for m in messages],
  60. "tools": [t.to_api_dict() for t in tools] if tools else None,
  61. "temperature": temp,
  62. }
  63. extra_body = self._extra_body()
  64. if extra_body:
  65. kwargs["extra_body"] = extra_body
  66. response = self._client.chat.completions.create(**kwargs)
  67. raw_message = response.choices[0].message
  68. result = self._parse_response(raw_message)
  69. if self.logger:
  70. self.logger.log_llm_output(iteration, result, raw_response=response)
  71. return result
  72. async def achat(
  73. self,
  74. messages: list[Message],
  75. tools: list[ToolDefinition] | None = None,
  76. temperature: float | None = None,
  77. *,
  78. iteration: int = 0,
  79. ) -> Message:
  80. """Async chat completion."""
  81. temp = temperature if temperature is not None else self.settings.agent_temperature
  82. if self.logger:
  83. self.logger.log_llm_input(iteration, self.model, messages, tools, temp)
  84. kwargs: dict[str, Any] = {
  85. "model": self.model,
  86. "messages": [m.to_api_dict() for m in messages],
  87. "tools": [t.to_api_dict() for t in tools] if tools else None,
  88. "temperature": temp,
  89. }
  90. extra_body = self._extra_body()
  91. if extra_body:
  92. kwargs["extra_body"] = extra_body
  93. response = await self._async_client.chat.completions.create(**kwargs)
  94. raw_message = response.choices[0].message
  95. result = self._parse_response(raw_message)
  96. if self.logger:
  97. self.logger.log_llm_output(iteration, result, raw_response=response)
  98. return result
  99. def stream(
  100. self,
  101. messages: list[Message],
  102. tools: list[ToolDefinition] | None = None,
  103. temperature: float | None = None,
  104. *,
  105. iteration: int = 0,
  106. ) -> Iterator[str]:
  107. """Stream text content chunks from the model."""
  108. temp = temperature if temperature is not None else self.settings.agent_temperature
  109. if self.logger:
  110. self.logger.log_llm_input(iteration, self.model, messages, tools, temp)
  111. kwargs: dict[str, Any] = {
  112. "model": self.model,
  113. "messages": [m.to_api_dict() for m in messages],
  114. "tools": [t.to_api_dict() for t in tools] if tools else None,
  115. "temperature": temp,
  116. "stream": True,
  117. }
  118. extra_body = self._extra_body()
  119. if extra_body:
  120. kwargs["extra_body"] = extra_body
  121. stream = self._client.chat.completions.create(**kwargs)
  122. chunks: list[str] = []
  123. for chunk in stream:
  124. delta = chunk.choices[0].delta
  125. if delta.content:
  126. chunks.append(delta.content)
  127. yield delta.content
  128. if self.logger:
  129. full_content = "".join(chunks)
  130. self.logger.log_llm_output(
  131. iteration,
  132. Message(role=Role.ASSISTANT, content=full_content),
  133. )
  134. async def astream(
  135. self,
  136. messages: list[Message],
  137. tools: list[ToolDefinition] | None = None,
  138. temperature: float | None = None,
  139. *,
  140. iteration: int = 0,
  141. ) -> AsyncIterator[str]:
  142. """Async stream text content chunks."""
  143. temp = temperature if temperature is not None else self.settings.agent_temperature
  144. if self.logger:
  145. self.logger.log_llm_input(iteration, self.model, messages, tools, temp)
  146. kwargs: dict[str, Any] = {
  147. "model": self.model,
  148. "messages": [m.to_api_dict() for m in messages],
  149. "tools": [t.to_api_dict() for t in tools] if tools else None,
  150. "temperature": temp,
  151. "stream": True,
  152. }
  153. extra_body = self._extra_body()
  154. if extra_body:
  155. kwargs["extra_body"] = extra_body
  156. stream = await self._async_client.chat.completions.create(**kwargs)
  157. chunks: list[str] = []
  158. async for chunk in stream:
  159. delta = chunk.choices[0].delta
  160. if delta.content:
  161. chunks.append(delta.content)
  162. yield delta.content
  163. if self.logger:
  164. full_content = "".join(chunks)
  165. self.logger.log_llm_output(
  166. iteration,
  167. Message(role=Role.ASSISTANT, content=full_content),
  168. )
  169. def _parse_response(self, choice_message: Any) -> Message:
  170. tool_calls = None
  171. if choice_message.tool_calls:
  172. tool_calls = [
  173. ToolCall(
  174. id=tc.id,
  175. name=tc.function.name,
  176. arguments=tc.function.arguments,
  177. )
  178. for tc in choice_message.tool_calls
  179. ]
  180. reasoning = getattr(choice_message, "reasoning", None)
  181. if reasoning is not None and not isinstance(reasoning, str):
  182. reasoning = str(reasoning)
  183. return Message(
  184. role=Role.ASSISTANT,
  185. content=choice_message.content,
  186. tool_calls=tool_calls,
  187. reasoning=reasoning,
  188. )
  189. def set_model(self, model: str) -> None:
  190. """Switch to a different model at runtime."""
  191. self.model = model