diff --git a/archytas/models/openai.py b/archytas/models/openai.py index 43887e0..39bac2f 100644 --- a/archytas/models/openai.py +++ b/archytas/models/openai.py @@ -100,6 +100,11 @@ def initialize_model(self, **kwargs): model=model, tiktoken_model_name=titoken_model_name, ) + # Optional custom endpoint (LiteLLM, an org gateway, or any other + # OpenAI-compatible server), configured per-provider like azure's + # `endpoint`. + if self.config.model_extra and self.config.model_extra.get("base_url"): + model_kwargs["base_url"] = self.config.model_extra["base_url"] reasoning = self._get_reasoning_config(model) if reasoning is not None: model_kwargs["reasoning"] = reasoning diff --git a/archytas/models/openrouter.py b/archytas/models/openrouter.py index b94fcfb..d18fe87 100644 --- a/archytas/models/openrouter.py +++ b/archytas/models/openrouter.py @@ -1,220 +1,96 @@ +"""OpenRouter provider, served over OpenRouter's OpenAI-compatible API. + +OpenRouter (https://openrouter.ai) exposes many providers' models behind the +OpenAI chat-completions API (including tool calling), so this provider is a +thin subclass of `OpenAIModel` pointed at the OpenRouter base URL rather than +a bespoke client. Model ids are namespaced by upstream provider, e.g. +`openai/gpt-4o-mini`, `anthropic/claude-3.5-sonnet`, `qwen/qwen3-coder:free`. +""" +import logging import os -import asyncio -import json -from typing import Any, Optional, Sequence, cast from functools import lru_cache -import logging -from toki import Model -from toki.openrouter import OpenRouterMessage, OpenRouterToolCall, OpenRouterToolFunction -from toki.openrouter_models import ModelName, Attr, attributes_map - -logger = logging.getLogger(__name__) - -from .base import BaseArchytasModel, ModelConfig -from ..exceptions import AuthenticationError, ExecutionError -from langchain_core.messages import BaseMessage, HumanMessage, SystemMessage, AIMessage, ToolCall as LangChainToolCall, ToolMessage, FunctionMessage - - -def to_langchain_tool_call(tool_call: OpenRouterToolCall) -> LangChainToolCall: - return LangChainToolCall(name=tool_call['function']['name'], args=json.loads(tool_call['function']['arguments'] or '{}'), id=tool_call['id'], type="tool_call") - -def to_openrouter_tool_call(tool_call: LangChainToolCall) -> OpenRouterToolCall: - return OpenRouterToolCall(id=tool_call['id'] or '', type="function", function=OpenRouterToolFunction(name=tool_call["name"], arguments=json.dumps(tool_call["args"]))) - - -class ChatOpenRouter: - """Minimal adapter that emulates a chat model interface over OpenRouter's REST API.""" - def __init__(self, *, model: str, api_key: str, base_url: str = "https://openrouter.ai/api/v1/chat/completions") -> None: - self.model = model - self.api_key = api_key - self.base_url = base_url - self._tools: Optional[Sequence[Any]] = None - - # Delegate to the minimal client - # If model isn't recognized in our Literal list, still construct with raw string - # this allows using newer models that haven't been updated in the openrouter_models.py file yet - self._client = Model(model=cast(ModelName, model), openrouter_api_key=api_key) # type: ignore[arg-type] - - # get the model attributes - try: - self._attrs = attributes_map[cast(ModelName, model)] - except KeyError: - self._attrs = Attr(context_size=200_000, supports_tools=True) - logger.warning(f"Unrecognized OpenRouter model: '{model}' (this implies model is not officially listed in toki.openrouter_models). To get most up-to-date models list, consider cutting a new toki release (after regenerating models list with `toki-fetch-models` command). Attempting to continue with the following Attributes: {self._attrs}") - - if not self._attrs.supports_tools: - raise ValueError(f"OpenRouter model '{model}' does not support tools. Archytas requires models to support tools. Please use a different model.") - - def bind_tools(self, tools: Sequence[Any]): - self._tools = tools - self._schemas = [] - for tool in tools: - langchain_schema: dict = tool.tool_call_schema.schema() - schema = { - 'type': 'function', - 'function': { - 'name': tool.name, - 'description': tool.description, - 'parameters': { - 'type': 'object', - 'properties': langchain_schema['properties'], - 'required': langchain_schema.get('required', []) - }, - } - } - self._schemas.append(schema) - - return self - - def _convert_messages(self, messages: list[BaseMessage]) -> list[OpenRouterMessage]: - def serialize_content(content: Any) -> str: - if isinstance(content, str): - return content - if isinstance(content, list): - parts: list[str] = [] - for item in content: - if isinstance(item, dict) and item.get("type") == "text" and "text" in item: - parts.append(item["text"]) - else: - parts.append(json.dumps(item)) - return "\n".join(parts) - return json.dumps(content) +from typing import Optional - converted: list[OpenRouterMessage] = [] - for msg in messages: - content = serialize_content(msg.content) - match msg: - case HumanMessage(): - converted.append({"role": "user", "content": content}) - case AIMessage(tool_calls=list() as tool_calls): - converted.append({"role": "assistant", 'content': content, 'tool_calls': list(map(to_openrouter_tool_call, tool_calls))}) - case AIMessage(): # Message without tool calls - converted.append({"role": "assistant", "content": content}) - case SystemMessage(): - converted.append({"role": "system", "content": content}) - case ToolMessage(): - converted.append({"role": "tool", "tool_call_id": msg.tool_call_id, "content": content}) - case _: - raise ValueError(f"Unexpected message type: {type(msg)}\n{msg=}") - return converted +import httpx +from langchain_openai.chat_models import ChatOpenAI - def get_num_tokens_from_messages(self, *, messages: list[BaseMessage], tools: Optional[Sequence[Any]] = None) -> int: - """Call OpenRouter to estimate prompt tokens by sending a completion request - that is configured to produce no output tokens. +from .openai import OpenAIModel +from ..exceptions import AuthenticationError - Returns the `prompt_tokens` reported in the response `usage`. - """ - raise NotImplementedError("Token count estimation for OpenRouter is not implemented") - - # TODO: this is very slow... basically it doubles the time to start generating a response - if True: #self.skip_token_count: - logger.warning("Skipping token count estimation for OpenRouter model") - return 0 - - # Convert messages to OpenRouter format - converted_messages = self._convert_messages(messages) - - # Build tool schemas if tools were provided; otherwise, reuse any bound schemas - schemas: list[dict] | None = None - if tools is not None: - try: - tmp_schemas: list[dict] = [] - for tool in tools: - langchain_schema: dict = tool.tool_call_schema.schema() - tmp_schemas.append({ - 'type': 'function', - 'function': { - 'name': tool.name, - 'description': tool.description, - 'parameters': langchain_schema['properties'], - 'required': langchain_schema['required'], - } - }) - schemas = tmp_schemas - except Exception: - # If tool introspection fails, ignore tools for token counting - schemas = None - else: - schemas = getattr(self, "_schemas", None) - - # Try with zero max tokens to avoid any generation. If the API rejects 0, - # fall back to 1 token with an immediate stop to minimize output tokens. - try: - self._client.complete(converted_messages, stream=False, tools=schemas, max_tokens=0) - except Exception as first_error: - try: - self._client.complete(converted_messages, stream=False, tools=schemas, max_tokens=1, stop=[""]) - except Exception as fallback_error: - raise ExecutionError( - f"Failed to estimate prompt tokens via OpenRouter. First attempt with max_tokens=0 error: {first_error}. " - f"Fallback with max_tokens=1 and immediate stop also failed: {fallback_error}" - ) from fallback_error - - usage = self._client._usage_metadata - if usage is None: - raise ExecutionError("OpenRouter did not return usage metadata for the token count request.") - return usage['prompt_tokens'] - - def invoke(self, input: list[BaseMessage], *args, **kwargs) -> AIMessage: - # Convert LangChain messages to OpenRouter format - converted_messages = self._convert_messages(input) - - response = self._client.complete(converted_messages, stream=False, tools=self._schemas, **kwargs) +logger = logging.getLogger(__name__) - assert self._client._usage_metadata is not None, "INTERNAL ERROR: Usage metadata was not set for previous completion call" - usage_metadata = { - 'input_tokens': self._client._usage_metadata['prompt_tokens'], - 'output_tokens': self._client._usage_metadata['completion_tokens'], - 'total_tokens': self._client._usage_metadata['total_tokens'] +OPENROUTER_BASE_URL = "https://openrouter.ai/api/v1" + +# Used when the model isn't found in the live model listing. Conservative on +# purpose: overestimating the window delays summarization until requests fail. +DEFAULT_CONTEXT_SIZE = 128_000 + +# Attribution headers, per OpenRouter's app-discoverability convention. +ATTRIBUTION_HEADERS = { + "HTTP-Referer": "https://github.com/jataware/archytas", + "X-Title": "Archytas", +} + + +@lru_cache(maxsize=1) +def _openrouter_context_lengths() -> dict[str, int]: + """Fetch context sizes for all models from OpenRouter's public model list. + + Cached for the process lifetime; returns an empty mapping on any failure so + callers fall back to DEFAULT_CONTEXT_SIZE rather than erroring. + """ + try: + response = httpx.get(f"{OPENROUTER_BASE_URL}/models", timeout=10) + response.raise_for_status() + models = response.json().get("data", []) + return { + model["id"]: int(model["context_length"]) + for model in models + if model.get("id") and model.get("context_length") } + except Exception as err: + logger.warning("Unable to fetch OpenRouter model list (%s); context sizes will use the default.", err) + return {} - if isinstance(response, dict): - return AIMessage(content=response.get('thought') or "", tool_calls=list(map(to_langchain_tool_call, response['tool_calls'])), usage_metadata=usage_metadata) - return AIMessage(content=response, usage_metadata=usage_metadata) - - async def ainvoke(self, input: list[BaseMessage], *args, **kwargs) -> AIMessage: - loop = asyncio.get_running_loop() - return await loop.run_in_executor(None, lambda: self.invoke(input, *args, **kwargs)) - - -class OpenRouterModel(BaseArchytasModel): - """Archytas backend model for OpenRouter using direct REST calls.""" - DEFAULT_MODEL = "openrouter/auto" - api_key: str = "" +class OpenRouterModel(OpenAIModel): + DEFAULT_MODEL = "openai/gpt-4o-mini" def auth(self, **kwargs) -> None: - self.api_key = ( + # Unlike OpenAIModel.auth, no environment round-trip: the key is used + # explicitly and must never be written to (or read back from) + # OPENAI_API_KEY, which may belong to a different provider entry. + api_key = ( kwargs.get("api_key") - or getattr(self.config, "api_key", None) - or os.getenv("OPENROUTER_API_KEY", "") + or self.config.api_key + or os.environ.get("OPENROUTER_API_KEY", "") + ) + if not api_key: + raise AuthenticationError( + "No OpenRouter API key found. Set one in the provider configuration or via OPENROUTER_API_KEY." + ) + self.config.api_key = api_key + + def initialize_model(self, **kwargs): + return ChatOpenAI( + model=self.config.model_name or self.DEFAULT_MODEL, + api_key=self.config.api_key, + base_url=OPENROUTER_BASE_URL, + # Namespaced ids like "anthropic/claude-3.5-sonnet" have no tiktoken + # encoding; pin a stable one so token estimates are quiet and + # consistent (they are budgeting approximations either way). + tiktoken_model_name="gpt-4o", + default_headers=ATTRIBUTION_HEADERS, ) - if not self.api_key: - raise AuthenticationError("No OpenRouter API Key found. Set OPENROUTER_API_KEY or pass api_key.") - - def initialize_model(self, **kwargs): # pyright: ignore[reportIncompatibleMethodOverride] - model_name = getattr(self.config, "model_name", None) or self.DEFAULT_MODEL - return ChatOpenRouter(model=str(model_name), api_key=self.api_key) - - async def get_num_tokens_from_messages( - self, - messages: "list[BaseMessage]", - tools: Optional[Sequence] = None, - ) -> int: - try: - return self._model.get_num_tokens_from_messages(messages=messages, tools=tools) - except Exception: - return 0 @lru_cache() def contextsize(self, model_name: Optional[str] = None) -> int | None: name = model_name or self.model_name - default_value = 200_000 - if model_name is not None: - try: - return attributes_map[cast(ModelName, model_name)].context_size - except KeyError: - pass - # Fallback default for safety so summarization threshold is usable - logger.warning(f"OpenRouter context size unknown for model '{name}' (this implies model is not officially listed in archytas/models/openrouter_models.py, i.e. consider regenerating the file with `create_models_types_file()`). Using default context size: {default_value}.") - return default_value \ No newline at end of file + context_size = _openrouter_context_lengths().get(name) + if context_size: + return context_size + logger.warning( + "OpenRouter context size unknown for model %r; using default of %d.", + name, DEFAULT_CONTEXT_SIZE, + ) + return DEFAULT_CONTEXT_SIZE diff --git a/pyproject.toml b/pyproject.toml index b38c15c..fde1c62 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -34,7 +34,6 @@ dependencies = [ "langchain-aws>=1.0", "azure-ai-inference>=1.0.0b9", "jinja2>=3.1.6", - "toki>=0.2.8", "langchain-mcp-adapters>=0.1.0", ] diff --git a/tests/test_openrouter.py b/tests/test_openrouter.py new file mode 100644 index 0000000..51a6d87 --- /dev/null +++ b/tests/test_openrouter.py @@ -0,0 +1,81 @@ +"""Tests for the OpenRouter provider (OpenAI-compatible subclass). + +No network calls: construction is enough to verify endpoint wiring, key +precedence, and env hygiene. Context-size lookup is exercised against a +patched model listing. +""" +import pytest + +import archytas.models.openrouter as openrouter_mod +from archytas.exceptions import AuthenticationError +from archytas.models.openai import OpenAIModel +from archytas.models.openrouter import OPENROUTER_BASE_URL, DEFAULT_CONTEXT_SIZE, OpenRouterModel + + +def _make(model_name="openai/gpt-4o-mini", **config): + return OpenRouterModel({"model_name": model_name, **config}) + + +class TestAuth: + def test_explicit_key_wins_over_env(self, monkeypatch): + monkeypatch.setenv("OPENROUTER_API_KEY", "sk-or-env-key") + model = _make(api_key="sk-or-config-key") + assert model.config.api_key == "sk-or-config-key" + + def test_falls_back_to_env(self, monkeypatch): + monkeypatch.setenv("OPENROUTER_API_KEY", "sk-or-env-key") + model = _make() + assert model.config.api_key == "sk-or-env-key" + + def test_no_key_raises(self, monkeypatch): + monkeypatch.delenv("OPENROUTER_API_KEY", raising=False) + with pytest.raises(AuthenticationError): + _make() + + def test_does_not_touch_openai_env(self, monkeypatch): + # The OpenRouter key must not leak into OPENAI_API_KEY (which may + # belong to a different provider entry in the same process). + monkeypatch.delenv("OPENAI_API_KEY", raising=False) + _make(api_key="sk-or-config-key") + import os + assert "OPENAI_API_KEY" not in os.environ + + +class TestClientWiring: + def test_points_at_openrouter(self): + model = _make(api_key="sk-or-key") + client = model._model + assert str(client.openai_api_base).startswith(OPENROUTER_BASE_URL) + assert client.model_name == "openai/gpt-4o-mini" + assert client.openai_api_key.get_secret_value() == "sk-or-key" + + +class TestContextSize: + def test_uses_live_listing(self, monkeypatch): + monkeypatch.setattr( + openrouter_mod, "_openrouter_context_lengths", + lambda: {"anthropic/claude-3.5-sonnet": 200_000}, + ) + model = _make(model_name="anthropic/claude-3.5-sonnet", api_key="sk-or-key") + assert model.contextsize("anthropic/claude-3.5-sonnet") == 200_000 + + def test_falls_back_to_default(self, monkeypatch): + monkeypatch.setattr(openrouter_mod, "_openrouter_context_lengths", lambda: {}) + model = _make(api_key="sk-or-key") + assert model.contextsize("someone/unknown-model") == DEFAULT_CONTEXT_SIZE + + +class TestOpenAIBaseUrlPassthrough: + def test_base_url_forwarded_from_config(self, monkeypatch): + monkeypatch.setenv("OPENAI_API_KEY", "sk-dummy") + model = OpenAIModel({ + "model_name": "gpt-4o-mini", + "api_key": "sk-dummy", + "base_url": "http://localhost:4000/v1", + }) + assert str(model._model.openai_api_base).startswith("http://localhost:4000/v1") + + def test_no_base_url_uses_default(self, monkeypatch): + monkeypatch.setenv("OPENAI_API_KEY", "sk-dummy") + model = OpenAIModel({"model_name": "gpt-4o-mini", "api_key": "sk-dummy"}) + assert model._model.openai_api_base in (None, "")