llm_client.py 2.0 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162
  1. from typing import AsyncIterator, Optional
  2. from openai import AsyncOpenAI
  3. from app.core.config import get_settings
  4. class LLMClient:
  5. def __init__(self):
  6. self._client = None
  7. def _ensure_client(self):
  8. if self._client is not None:
  9. return
  10. settings = get_settings()
  11. is_local = settings.qwen_base_url.startswith("http://localhost")
  12. api_key = "local" if is_local else settings.qwen_api_key
  13. if not api_key and not is_local:
  14. raise ValueError("QWEN_API_KEY 未设置,请在 .env 文件中填入 API Key")
  15. self._client = AsyncOpenAI(
  16. api_key=api_key,
  17. base_url=settings.qwen_base_url,
  18. )
  19. self.model = settings.qwen_model
  20. async def chat(
  21. self,
  22. messages: list[dict],
  23. temperature: Optional[float] = None,
  24. max_tokens: Optional[int] = None,
  25. ) -> str:
  26. self._ensure_client()
  27. settings = get_settings()
  28. response = await self._client.chat.completions.create(
  29. model=self.model,
  30. messages=messages,
  31. temperature=temperature or settings.qwen_temperature,
  32. max_tokens=max_tokens or settings.qwen_max_tokens,
  33. )
  34. return response.choices[0].message.content or ""
  35. async def chat_stream(
  36. self,
  37. messages: list[dict],
  38. temperature: Optional[float] = None,
  39. max_tokens: Optional[int] = None,
  40. ) -> AsyncIterator[str]:
  41. self._ensure_client()
  42. settings = get_settings()
  43. stream = await self._client.chat.completions.create(
  44. model=self.model,
  45. messages=messages,
  46. temperature=temperature or settings.qwen_temperature,
  47. max_tokens=max_tokens or settings.qwen_max_tokens,
  48. stream=True,
  49. )
  50. async for chunk in stream:
  51. if chunk.choices and chunk.choices[0].delta.content:
  52. yield chunk.choices[0].delta.content
  53. llm_client = LLMClient()