Spaces:
Paused
Paused
| """LLM provider abstraction for OpenAI, Anthropic, and OpenRouter.""" | |
| from __future__ import annotations | |
| import logging | |
| from abc import ABC, abstractmethod | |
| from typing import Any | |
| from hermes.config.settings import get_settings | |
| logger = logging.getLogger(__name__) | |
| class LLMProvider(ABC): | |
| """Abstract LLM provider.""" | |
| async def chat( | |
| self, messages: list[dict[str, str]], temperature: float | None = None, max_tokens: int | None = None | |
| ) -> str: | |
| """Send chat completion request.""" | |
| async def chat_stream( | |
| self, messages: list[dict[str, str]], temperature: float | None = None, max_tokens: int | None = None | |
| ): | |
| """Stream chat completion response. Yields chunks of text.""" | |
| # Default: yield the full response as one chunk | |
| result = await self.chat(messages, temperature, max_tokens) | |
| yield result | |
| def count_tokens(self, text: str) -> int: | |
| """Count tokens in text.""" | |
| class OpenAIProvider(LLMProvider): | |
| """OpenAI chat completion provider.""" | |
| def __init__(self, api_key: str, model: str = "gpt-4o", base_url: str | None = None) -> None: | |
| self.api_key = api_key | |
| self.model = model | |
| self.base_url = base_url | |
| self._client: Any = None | |
| async def _get_client(self) -> Any: | |
| if self._client is None: | |
| from openai import AsyncOpenAI | |
| self._client = AsyncOpenAI(api_key=self.api_key, base_url=self.base_url) | |
| return self._client | |
| async def chat( | |
| self, messages: list[dict[str, str]], temperature: float | None = None, max_tokens: int | None = None | |
| ) -> str: | |
| client = await self._get_client() | |
| import asyncio | |
| from openai import RateLimitError | |
| last_error: Exception | None = None | |
| for attempt in range(3): | |
| try: | |
| response = await asyncio.wait_for( | |
| client.chat.completions.create( | |
| model=self.model, | |
| messages=messages, | |
| temperature=temperature or 0.1, | |
| max_tokens=max_tokens or 4096, | |
| ), | |
| timeout=110.0, | |
| ) | |
| return response.choices[0].message.content or "" | |
| except RateLimitError as e: | |
| last_error = e | |
| wait = 2 ** attempt * 10 | |
| logger.warning(f"Rate limited, retrying in {wait}s (attempt {attempt + 1}/3)") | |
| await asyncio.sleep(wait) | |
| except Exception as e: | |
| msg = str(e) | |
| for key in (self.api_key, self.model): | |
| if key and key in msg: | |
| msg = msg.replace(key, "***") | |
| logger.error(f"OpenAI API error: {msg}") | |
| raise RuntimeError("LLM API call failed") from e | |
| if last_error is not None: | |
| raise RuntimeError("LLM API rate limit exceeded after 3 retries") from last_error | |
| return "" | |
| def count_tokens(self, text: str) -> int: | |
| try: | |
| import tiktoken | |
| enc = tiktoken.encoding_for_model(self.model) | |
| return len(enc.encode(text)) | |
| except Exception: | |
| return len(text) // 4 | |
| class AnthropicProvider(LLMProvider): | |
| """Anthropic Claude chat completion provider.""" | |
| def __init__(self, api_key: str, model: str = "claude-sonnet-4-20250514") -> None: | |
| self.api_key = api_key | |
| self.model = model | |
| self._client: Any = None | |
| async def _get_client(self) -> Any: | |
| if self._client is None: | |
| from anthropic import AsyncAnthropic | |
| self._client = AsyncAnthropic(api_key=self.api_key) | |
| return self._client | |
| async def chat( | |
| self, messages: list[dict[str, str]], temperature: float | None = None, max_tokens: int | None = None | |
| ) -> str: | |
| client = await self._get_client() | |
| system_msgs = [m for m in messages if m["role"] == "system"] | |
| chat_msgs = [m for m in messages if m["role"] != "system"] | |
| kwargs: dict[str, Any] = { | |
| "model": self.model, | |
| "max_tokens": max_tokens or 4096, | |
| "temperature": temperature or 0.1, | |
| } | |
| if system_msgs: | |
| kwargs["system"] = system_msgs[-1]["content"] | |
| if chat_msgs: | |
| kwargs["messages"] = chat_msgs | |
| response = await client.messages.create(**kwargs) | |
| return response.content[0].text if response.content else "" | |
| def count_tokens(self, text: str) -> int: | |
| try: | |
| import tiktoken | |
| enc = tiktoken.encoding_for_model("cl100k_base") | |
| return len(enc.encode(text)) | |
| except Exception: | |
| return len(text) // 4 | |
| class OpenRouterProvider(LLMProvider): | |
| """OpenRouter chat completion provider. | |
| Uses OpenAI-compatible API to access 200+ models including | |
| GPT, Claude, Gemini, DeepSeek, Llama, Mistral, Qwen, and more. | |
| """ | |
| def __init__( | |
| self, | |
| api_key: str, | |
| model: str = "openai/gpt-4o", | |
| base_url: str = "https://openrouter.ai/api/v1", | |
| site_url: str = "", | |
| app_name: str = "Hermes Platform", | |
| ) -> None: | |
| self.api_key = api_key | |
| self.model = model | |
| self.base_url = base_url | |
| self.site_url = site_url | |
| self.app_name = app_name | |
| self._client: Any = None | |
| async def _get_client(self) -> Any: | |
| if self._client is None: | |
| from openai import AsyncOpenAI | |
| headers: dict[str, str] = { | |
| "Authorization": f"Bearer {self.api_key}", | |
| } | |
| if self.site_url: | |
| headers["HTTP-Referer"] = self.site_url | |
| if self.app_name: | |
| headers["X-Title"] = self.app_name | |
| import httpx | |
| self._client = AsyncOpenAI( | |
| api_key=self.api_key, | |
| base_url=self.base_url, | |
| default_headers=headers, | |
| http_client=httpx.AsyncClient(timeout=httpx.Timeout(120.0)), | |
| ) | |
| return self._client | |
| async def chat( | |
| self, messages: list[dict[str, str]], temperature: float | None = None, max_tokens: int | None = None | |
| ) -> str: | |
| client = await self._get_client() | |
| import asyncio | |
| from openai import RateLimitError | |
| last_error: Exception | None = None | |
| for attempt in range(3): | |
| try: | |
| response = await asyncio.wait_for( | |
| client.chat.completions.create( | |
| model=self.model, | |
| messages=messages, | |
| temperature=temperature or 0.1, | |
| max_tokens=max_tokens or 4096, | |
| ), | |
| timeout=110.0, | |
| ) | |
| return response.choices[0].message.content or "" | |
| except RateLimitError as e: | |
| last_error = e | |
| wait = 2 ** attempt * 10 | |
| logger.warning(f"Rate limited, retrying in {wait}s (attempt {attempt + 1}/3)") | |
| await asyncio.sleep(wait) | |
| except Exception as e: | |
| msg = str(e) | |
| for key in (self.api_key, self.model): | |
| if key and key in msg: | |
| msg = msg.replace(key, "***") | |
| logger.error(f"OpenRouter API error: {msg}") | |
| raise RuntimeError("LLM API call failed") from e | |
| if last_error is not None: | |
| raise RuntimeError("LLM API rate limit exceeded after 3 retries") from last_error | |
| return "" | |
| async def chat_stream( | |
| self, messages: list[dict[str, str]], temperature: float | None = None, max_tokens: int | None = None | |
| ): | |
| """Stream chat completion response.""" | |
| client = await self._get_client() | |
| try: | |
| stream = await client.chat.completions.create( | |
| model=self.model, | |
| messages=messages, | |
| temperature=temperature or 0.1, | |
| max_tokens=max_tokens or 4096, | |
| stream=True, | |
| ) | |
| async for chunk in stream: | |
| if chunk.choices and chunk.choices[0].delta.content: | |
| yield chunk.choices[0].delta.content | |
| except Exception as e: | |
| msg = str(e) | |
| for key in (self.api_key, self.model): | |
| if key and key in msg: | |
| msg = msg.replace(key, "***") | |
| logger.error(f"OpenRouter stream error: {msg}") | |
| raise RuntimeError("LLM streaming failed") from e | |
| def count_tokens(self, text: str) -> int: | |
| try: | |
| import tiktoken | |
| enc = tiktoken.encoding_for_model("cl100k_base") | |
| return len(enc.encode(text)) | |
| except Exception: | |
| return len(text) // 4 | |
| class MockProvider(LLMProvider): | |
| """Mock provider for testing and development.""" | |
| def __init__(self) -> None: | |
| self.call_count = 0 | |
| async def chat( | |
| self, messages: list[dict[str, str]], temperature: float | None = None, max_tokens: int | None = None | |
| ) -> str: | |
| self.call_count += 1 | |
| last = messages[-1]["content"] if messages else "" | |
| return f"[Mock LLM Response to: {last[:80]}...]" | |
| def count_tokens(self, text: str) -> int: | |
| return len(text) // 4 | |
| def get_llm_provider() -> LLMProvider: | |
| """Get the configured LLM provider.""" | |
| settings = get_settings() | |
| if settings.model.provider == "openrouter": | |
| if not settings.model.openrouter_api_key: | |
| logger.warning("OpenRouter configured but no API key set, using mock") | |
| return MockProvider() | |
| return OpenRouterProvider( | |
| api_key=settings.model.openrouter_api_key, | |
| model=settings.model.name, | |
| base_url=settings.model.openrouter_base_url, | |
| site_url=settings.model.openrouter_site_url, | |
| app_name=settings.model.openrouter_app_name, | |
| ) | |
| if settings.model.provider == "opencode": | |
| if not settings.model.api_key: | |
| logger.warning("OpenCode configured but no API key set, using mock") | |
| return MockProvider() | |
| return OpenAIProvider( | |
| api_key=settings.model.api_key, | |
| model=settings.model.name, | |
| base_url=settings.model.base_url or "https://opencode.ai/zen/go/v1", | |
| ) | |
| if not settings.model.api_key: | |
| logger.warning("No API key configured, using mock LLM provider") | |
| return MockProvider() | |
| if settings.model.provider == "anthropic": | |
| return AnthropicProvider( | |
| api_key=settings.model.api_key, | |
| model=settings.model.name, | |
| ) | |
| return OpenAIProvider( | |
| api_key=settings.model.api_key, | |
| model=settings.model.name, | |
| base_url=settings.model.base_url or None, | |
| ) | |
| class ObservedLLMProvider(LLMProvider): | |
| """Wrapper that adds OTel tracing + metrics to any LLM provider.""" | |
| def __init__(self, inner: LLMProvider) -> None: | |
| self._inner = inner | |
| async def chat( | |
| self, messages: list[dict[str, str]], temperature: float | None = None, max_tokens: int | None = None | |
| ) -> str: | |
| from hermes.observability.metrics import metrics as m | |
| m.increment("llm.calls.total") | |
| m.start_timer("llm.chat") | |
| try: | |
| result = await self._inner.chat(messages, temperature, max_tokens) | |
| m.stop_timer("llm.chat") | |
| m.increment("llm.calls.success") | |
| return result | |
| except Exception: | |
| m.stop_timer("llm.chat") | |
| m.increment("llm.calls.failure") | |
| raise | |
| async def chat_stream( | |
| self, messages: list[dict[str, str]], temperature: float | None = None, max_tokens: int | None = None | |
| ): | |
| """Stream with metrics.""" | |
| from hermes.observability.metrics import metrics as m | |
| m.increment("llm.stream.total") | |
| async for chunk in self._inner.chat_stream(messages, temperature, max_tokens): | |
| yield chunk | |
| m.increment("llm.stream.success") | |
| def count_tokens(self, text: str) -> int: | |
| return self._inner.count_tokens(text) | |