diff --git a/artemis/agents/flash/runner.py b/artemis/agents/flash/runner.py index 3ea3f686..d90ea556 100644 --- a/artemis/agents/flash/runner.py +++ b/artemis/agents/flash/runner.py @@ -165,6 +165,7 @@ def __init__(self, ctx: ArtemisContext, goal: str, max_turns: int | None = None) ctx, model_name=self.step_summarizer_cfg.model, retry_limit=self.memory_runtime_cfg.retry_limit, + model_provider=self.step_summarizer_cfg.provider, max_concurrency=self.memory_runtime_cfg.max_concurrency, flush_timeout_s=self.memory_runtime_cfg.flush_timeout_s, ) diff --git a/artemis/agents/flash/summarizer.py b/artemis/agents/flash/summarizer.py index 5b8a3cf9..7852a5ab 100644 --- a/artemis/agents/flash/summarizer.py +++ b/artemis/agents/flash/summarizer.py @@ -36,7 +36,7 @@ from artemis.context import ArtemisContext from artemis.memory.step_memory import JobKey, StepMemoryService -from artemis.services.llm import RobustChatModelWrapper, get_google_llm, get_llm +from artemis.services.llm import RobustChatModelWrapper, get_lens_llm from artemis.services.token_meter import record_llm_usage from artemis.utils.task_tree import format_actions_clean from artemis.utils.visualization import draw_action_overlay_on_image @@ -151,6 +151,7 @@ def __init__( model_name: str | None = None, retry_limit: int = 3, *, + model_provider: str | None = None, max_concurrency: int = 1, flush_timeout_s: float = 30.0, ): @@ -161,16 +162,16 @@ def __init__( flush_timeout_s=flush_timeout_s, ) - # Initialize lightweight VLM: prioritize explicit model_name + # Initialize lightweight VLM: prioritize explicit model_name. Routing + # is provider-aware: an explicit provider knob wins, Gemini names stay + # on Google, and anything else inherits the summarizer node's provider + # from the LLM config (see get_lens_llm). Construction errors surface + # instead of silently swapping in a different provider's model — a + # silent fallback is what previously hid provider misroutes behind + # per-step 404s. target_model = model_name or "gemini-2.5-flash-lite" self._model_name = target_model - try: - if model_name: - self._llm = get_google_llm(model_name=target_model, temperature=0.0) - else: - self._llm = get_llm(ctx, name="summarizer", is_utils=True) - except Exception: - self._llm = get_google_llm(model_name=target_model, temperature=0.0) + self._llm = get_lens_llm(ctx, target_model, model_provider, temperature=0.0) try: configured = getattr(self._llm, "model", None) or getattr(self._llm, "model_name", None) if isinstance(configured, str) and configured: diff --git a/artemis/config/agent.py b/artemis/config/agent.py index 6fb00d39..80ccadf4 100644 --- a/artemis/config/agent.py +++ b/artemis/config/agent.py @@ -18,7 +18,7 @@ from pathlib import Path from typing import Any, Literal -from pydantic import BaseModel, Field, model_validator +from pydantic import BaseModel, Field, field_validator, model_validator from artemis.config.constants import ( AGENT_CONFIG_FILENAME, @@ -27,6 +27,7 @@ ) from artemis.config.paths import ROOT_DIR, get_config_path from artemis.llm.google import VideoProcessing +from artemis.llm.router import ModelProvider from third_party.mobile_use.utils.file import load_jsonc from third_party.mobile_use.utils.logger import get_logger @@ -392,7 +393,21 @@ class StepSummarizerConfig(BaseModel): ) model: str = Field( default="gemini-2.5-flash-lite", - description="Lightweight model used for background step state summarization.", + description=( + "Lightweight model used for background step state summarization." + " Gemini names route to the Google provider automatically; any other" + " model ID needs 'provider' set explicitly or inherits the LLM" + " config's summarizer node provider (itself inheriting 'default')." + ), + ) + provider: str | None = Field( + default=None, + description=( + "Explicit provider for the step summarizer model ('google'," + " 'custom', 'openai', ...). None (default) auto-resolves: Gemini" + " model names use Google; other IDs inherit the summarizer node's" + " provider from the LLM config." + ), ) prune_history_xml: bool = Field( default=True, @@ -409,6 +424,14 @@ class StepSummarizerConfig(BaseModel): ) model_config = {"extra": "allow"} + @field_validator("provider") + @classmethod + def _validate_provider(cls, v: str | None) -> str | None: + """Fail fast on unknown provider names at config-load time.""" + if v is not None and str(v).strip(): + ModelProvider.from_string(v) + return v + class MemoryRuntimeConfig(BaseModel): """Scheduling options for the shared step-memory runtime (agent.memory.runtime).""" @@ -689,7 +712,23 @@ class MemoryChunkingConfig(BaseModel): ) model: str = Field( default="gemini-3.8-flash", - description="Model used for the chunk-level StepCapsuleLens (bands ①+②).", + description=( + "Model used for the chunk-level StepCapsuleLens (bands ①+②)." + " Gemini names route to the Google provider automatically; any" + " other model ID needs 'provider' set explicitly or inherits the" + " LLM config's summarizer node provider (itself inheriting" + " 'default')." + ), + ) + provider: str | None = Field( + default=None, + description=( + "Explicit provider for the chunk capsule model ('google'," + " 'custom', 'openai', ...). None (default) auto-resolves: Gemini" + " model names use Google; other IDs inherit the summarizer node's" + " provider from the LLM config. Also applies to the capsule" + " fallback model resolved from the summarizer node's fallback." + ), ) max_chunks: int = Field( default=8, @@ -719,6 +758,14 @@ def _min_steps_within_cap(self) -> "MemoryChunkingConfig": ) return self + @field_validator("provider") + @classmethod + def _validate_provider(cls, v: str | None) -> str | None: + """Fail fast on unknown provider names at config-load time.""" + if v is not None and str(v).strip(): + ModelProvider.from_string(v) + return v + model_config = {"extra": "allow"} diff --git a/artemis/memory/__init__.py b/artemis/memory/__init__.py index d0f2f2e7..27b696a2 100644 --- a/artemis/memory/__init__.py +++ b/artemis/memory/__init__.py @@ -49,11 +49,13 @@ def ensure_step_memory(ctx): kwargs: dict = {} model_name = None + model_provider = None try: from artemis.config import load_agent_config cfg = load_agent_config() model_name = cfg.flash.step_summarizer.model + model_provider = cfg.flash.step_summarizer.provider kwargs = { "retry_limit": cfg.memory.runtime.retry_limit, "max_concurrency": cfg.memory.runtime.max_concurrency, @@ -65,7 +67,9 @@ def ensure_step_memory(ctx): exc_info=True, ) - service = VisualStepSummarizer(ctx, model_name=model_name, **kwargs) + service = VisualStepSummarizer( + ctx, model_name=model_name, model_provider=model_provider, **kwargs + ) try: ctx.step_memory = service except (AttributeError, TypeError, ValueError): diff --git a/artemis/memory/chunking.py b/artemis/memory/chunking.py index c0c91ea0..aec2526e 100644 --- a/artemis/memory/chunking.py +++ b/artemis/memory/chunking.py @@ -45,7 +45,6 @@ from langchain_core.messages import BaseMessage, HumanMessage, SystemMessage -from artemis.llm.google import is_google_provider from artemis.memory.step_memory import JobKey, StepLens, StepMemoryService from artemis.memory.transcript import format_session_offset from third_party.mobile_use.utils.logger import get_logger @@ -295,10 +294,13 @@ def __init__( llm: Any | None = None, *, ctx: Any = None, + model_provider: str | None = None, fallback_model_name: str | None = None, + fallback_model_provider: str | None = None, fallback_llm: Any | None = None, ): self._model_name = model_name or "gemini-3.8-flash" + self._model_provider = model_provider self._llm = llm self._ctx = ctx # Availability hardening: `chunking.model` is a dedicated model with no @@ -309,6 +311,7 @@ def __init__( self._fallback_model_name = ( fallback_model_name if fallback_model_name != self._model_name else None ) + self._fallback_model_provider = fallback_model_provider self._fallback_llm = fallback_llm try: self._prompt = self._PROMPT_PATH.read_text(encoding="utf-8") @@ -323,17 +326,22 @@ def __init__( def _get_llm(self): if self._llm is None: - from artemis.services.llm import get_google_llm + from artemis.services.llm import get_lens_llm - self._llm = get_google_llm(model_name=self._model_name, temperature=0.0) + self._llm = get_lens_llm( + self._ctx, self._model_name, self._model_provider, temperature=0.0 + ) return self._llm def _get_fallback_llm(self): if self._fallback_llm is None and self._fallback_model_name: - from artemis.services.llm import get_google_llm + from artemis.services.llm import get_lens_llm - self._fallback_llm = get_google_llm( - model_name=self._fallback_model_name, temperature=0.0 + self._fallback_llm = get_lens_llm( + self._ctx, + self._fallback_model_name, + self._fallback_model_provider, + temperature=0.0, ) return self._fallback_llm @@ -892,6 +900,7 @@ def __init__( self._min_steps = max(1, min(self._max_steps, int(getattr(cc, "min_steps", 3) or 3))) self._target_source_tokens = int(getattr(cc, "target_source_tokens", 2000) or 2000) self._model_name = getattr(cc, "model", None) or "gemini-3.8-flash" + self._model_provider = getattr(cc, "provider", None) self._max_chunks = int(getattr(cc, "max_chunks", 8) or 8) # None uses max_chunks as the era cap. self._max_eras = int(getattr(cc, "max_eras", None) or self._max_chunks) @@ -942,19 +951,29 @@ def _build_capsule_service(self, ctx: Any) -> StepMemoryService: f"Memory runtime config unavailable; using capsule service defaults: {exc}", exc_info=True, ) + fallback_model, fallback_provider = self._resolve_capsule_fallback(ctx) lens = StepCapsuleLens( self._model_name, ctx=ctx, - fallback_model_name=self._resolve_capsule_fallback_model(ctx), + model_provider=self._model_provider, + fallback_model_name=fallback_model, + # A fallback block without its own provider rides the chunking + # 'provider' knob; resolve eagerly so the lens stores what it + # will actually use. + fallback_model_provider=fallback_provider or self._model_provider, ) return ChunkCapsuleService(ctx, lens, **kwargs) - def _resolve_capsule_fallback_model(self, ctx: Any) -> str | None: - """Fallback model for capsule generation when `chunking.model` is down. + def _resolve_capsule_fallback(self, ctx: Any) -> tuple[str | None, str | None]: + """Fallback model + provider for capsule generation when `chunking.model` is down. Resolved from the LLM config's summarizer role (which inherits the - global default fallback unless overridden). Only same-provider (google) - fallbacks apply — the capsule lens rides the raw google model path. + global default fallback unless overridden). The fallback block carries + its own provider — that is how every other node's fallback resolves + (see ``_resolve_endpoint(use_fallback=True)``) — so it wins over the + chunking 'provider' knob. Any provider applies: the capsule lens rides + the provider-aware lens resolver, so a non-google summarizer fallback + (e.g. a gateway model) is now usable too. """ try: llm_cfg = getattr(ctx, "llm_config", None) if ctx is not None else None @@ -963,13 +982,13 @@ def _resolve_capsule_fallback_model(self, ctx: Any) -> str | None: llm_cfg = get_default_llm_config() fallback = getattr(getattr(llm_cfg, "summarizer", None), "fallback", None) - provider = str(getattr(fallback, "provider", "") or "") model = getattr(fallback, "model", None) - if model and is_google_provider(provider) and model != self._model_name: - return str(model) + provider = getattr(fallback, "provider", None) + if model and model != self._model_name: + return str(model), (str(provider) if provider else None) except Exception as exc: logger.debug(f"Capsule fallback model resolution skipped: {exc}", exc_info=True) - return None + return None, None @property def capsule_service(self) -> StepMemoryService: diff --git a/artemis/services/llm.py b/artemis/services/llm.py index ec91e6fa..470de3bb 100644 --- a/artemis/services/llm.py +++ b/artemis/services/llm.py @@ -43,6 +43,7 @@ from artemis.context import ArtemisContext from artemis.data_engine.trace import CURRENT_TRACE_ID, DataEngineCallbackHandler from artemis.llm.google import is_google_chat_model, is_google_provider +from artemis.llm.google.provider import is_gemini_model from artemis.llm.reliability import ( CircuitBreaker, FailureCategory, @@ -893,6 +894,104 @@ def get_google_llm( return ModelFactory.create_model(ep) +def get_lens_llm( + ctx: ArtemisContext | None, + model_name: str, + provider: str | None = None, + *, + temperature: float = 0.0, + timeout: float = 60.0, +) -> BaseChatModel: + """Provider-aware raw-model factory for step-memory lens calls. + + Resolves the provider for an explicitly configured lens model (step + summarizer, chunk-capsule lens) and returns the *raw* chat model — + deliberately not wrapped in :class:`RobustChatModelWrapper`, because the + lens services own their bounded retry/timeout loops and meter raw calls + explicitly (``lens:*`` sources). + + Resolution order: + 1. Explicit ``provider`` knob (new config field) — validated via + :meth:`ModelProvider.from_string`, so typos fail fast with the + standard "Unknown LLM provider" error. + 2. Gemini model names (optionally ``google/``- or ``gemini/``-prefixed) + route to GOOGLE. This keeps every pre-existing config byte-for-byte + compatible: a bare ``gemini-*`` name never needs a provider knob. + 3. Anything else inherits the provider from the LLM config's + ``summarizer`` node (which itself inherits the global ``default`` + block). Namespaced model IDs such as ``tensorx/deepseek/...`` are + passed through verbatim — never parsed as provider prefixes. + 4. On any resolution failure, fall back to CUSTOM (the OpenAI-compatible + gateway), matching how utils nodes behave when no config is live. + """ + # 1. Explicit provider knob wins and must be a known provider. + explicit = str(provider or "").strip() + if explicit: + resolved = ModelProvider.from_string(explicit) + llm_logger.info( + f"get_lens_llm: provider override {resolved.value!r} for model {model_name!r}" + ) + else: + # 2. Gemini names — bare, or google/gemini-prefixed — stay on GOOGLE. + bare = str(model_name or "").strip() + prefix, sep, rest = bare.partition("/") + if is_gemini_model(bare) or (sep and rest and prefix.lower() in ("google", "gemini")): + resolved = ModelProvider.GOOGLE + else: + # 3. Inherit from the summarizer node's config (-> default block). + resolved = _inherit_lens_provider(ctx) + _warn_if_non_google_routes_to_google(resolved, model_name) + ep = ModelEndpoint( + provider=resolved, + model_name=model_name, + temperature=temperature, + timeout_seconds=timeout, + ) + llm_logger.info(f"get_lens_llm: resolved {resolved.value}:{model_name}") + return ModelFactory.create_model(ep) + + +def _inherit_lens_provider(ctx: ArtemisContext | None) -> ModelProvider: + """Inherit the lens provider from the summarizer node's LLM config. + + ``summarizer`` is the natural donor: it inherits the global ``default`` + block exactly like every other node, and the step-memory runtime already + reads the summarizer fallback for capsule availability (see + ``_resolve_capsule_fallback_model``). Unknown/namespaced model IDs inherit + the donor's provider verbatim — gateway namespaces like ``tensorx/...`` + are never parsed as providers here. + """ + try: + llm_cfg = getattr(ctx, "llm_config", None) if ctx is not None else None + if llm_cfg is None: + from artemis.config.llm import get_default_llm_config as _get_default + + llm_cfg = _get_default() + donor = getattr(llm_cfg, "summarizer", None) + provider_val = getattr(donor, "provider", None) + if provider_val: + return ModelProvider.from_string(provider_val) + except Exception as exc: + llm_logger.debug(f"Lens provider inheritance failed; defaulting to custom: {exc}") + return ModelProvider.CUSTOM + + +def _warn_if_non_google_routes_to_google(provider: ModelProvider, model_name: str) -> None: + """Diagnosability guard: flag Gemini-API routing of non-Gemini names. + + The historical step-summarizer 404 came from a non-Google model ID being + sent verbatim to the Gemini Developer API; the request URL was fine, so + nothing else in the stack could have flagged it. Warn loudly when the + resolved provider would repeat that pattern. + """ + if provider == ModelProvider.GOOGLE and not is_gemini_model(model_name): + llm_logger.warning( + f"get_lens_llm: non-Gemini model {model_name!r} is routed to the Google " + "Gemini API and will likely 404; set an explicit 'provider' or inherit a " + "non-google default provider for this model." + ) + + def _resolve_endpoint( ctx: ArtemisContext, name: str, diff --git a/tests/unit/agents/test_flash_runner.py b/tests/unit/agents/test_flash_runner.py index f188b5e1..cbfd79cb 100644 --- a/tests/unit/agents/test_flash_runner.py +++ b/tests/unit/agents/test_flash_runner.py @@ -467,6 +467,38 @@ async def test_visual_lens_receives_the_recorded_action_shape(mock_context): ) +def test_flash_runner_forwards_step_summarizer_provider(mock_context, monkeypatch): + """The configured step_summarizer 'provider' knob reaches the lens.""" + from artemis.config.agent import AgentGlobalConfig + from artemis.agents.flash.summarizer import VisualStepSummarizer + + captured = {} + + # Patch at the runner's import site: the runner references the class + # directly, so intercept the construction kwargs wholesale. + monkeypatch.setattr( + "artemis.agents.flash.runner.VisualStepSummarizer", + lambda ctx, **kwargs: captured.update(kwargs) or Mock(), + ) + cfg = AgentGlobalConfig.model_validate( + { + "flash": { + "step_summarizer": { + "model": "tensorx/deepseek/deepseek-v4-flash", + "provider": "custom", + } + } + } + ) + monkeypatch.setattr("artemis.agents.flash.runner.load_agent_config", lambda: cfg, raising=False) + + with patch("artemis.controllers.unified_controller.get_driver"): + FlashRunner(mock_context, goal="g") + + assert captured["model_name"] == "tensorx/deepseek/deepseek-v4-flash" + assert captured["model_provider"] == "custom" + + # --- Native thinking (thought summaries) reach the step record ---------------------- diff --git a/tests/unit/memory/test_history_chunking.py b/tests/unit/memory/test_history_chunking.py index f24c6128..ff78c750 100644 --- a/tests/unit/memory/test_history_chunking.py +++ b/tests/unit/memory/test_history_chunking.py @@ -30,10 +30,12 @@ import json import re from types import SimpleNamespace +from unittest.mock import patch import pytest from langchain_core.messages import AIMessage, HumanMessage, ToolMessage +import artemis.memory.chunking as chunking_module from artemis.memory.chunking import ( CHUNK_PENDING_NOTE, ChunkState, @@ -1005,7 +1007,7 @@ async def ainvoke(self, messages): assert service.get_summary("chunk:1-2") is None -def test_capsule_fallback_model_resolution_google_only_and_not_primary(): +def test_capsule_fallback_resolution_provider_aware(): mgr = HistoryChunkManager( capsule_service=StubCapsuleService(), chunking_config=SimpleNamespace( @@ -1023,14 +1025,91 @@ def ctx_with(provider, model): ) ) - assert ( - mgr._resolve_capsule_fallback_model(ctx_with("google", "gemini-3.6-flash")) - == "gemini-3.6-flash" + assert mgr._resolve_capsule_fallback(ctx_with("google", "gemini-3.6-flash")) == ( + "gemini-3.6-flash", + "google", + ) + # Non-google fallbacks ride the provider-aware lens resolver now. + assert mgr._resolve_capsule_fallback(ctx_with("openai", "gpt-4o-mini")) == ( + "gpt-4o-mini", + "openai", ) - # Non-google fallbacks cannot ride the raw google model path. - assert mgr._resolve_capsule_fallback_model(ctx_with("openai", "gpt-4o-mini")) is None # A fallback identical to the primary adds nothing. - assert mgr._resolve_capsule_fallback_model(ctx_with("google", "gemini-3.7-flash")) is None + assert mgr._resolve_capsule_fallback(ctx_with("google", "gemini-3.7-flash")) == (None, None) + + +def test_capsule_fallback_provider_knob_applies_to_fallback_provider(): + """When the summarizer fallback block carries no provider of its own, the + chunking 'provider' knob applies to the fallback too.""" + captured = {} + + class RecordingLens(StepCapsuleLens): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + captured.update( + model_provider=self._model_provider, + fallback_model_name=self._fallback_model_name, + fallback_model_provider=self._fallback_model_provider, + ) + + # Fallback block without a provider: the knob decides. + ctx = SimpleNamespace( + llm_config=SimpleNamespace( + summarizer=SimpleNamespace(fallback=SimpleNamespace(model="vendor/fallback")) + ) + ) + mgr = HistoryChunkManager( + chunking_config=SimpleNamespace( + max_steps=12, + target_source_tokens=2000, + model="tensorx/deepseek/deepseek-v4-flash", + provider="custom", + max_chunks=8, + ), + ) + with patch.object(chunking_module, "StepCapsuleLens", RecordingLens): + mgr._build_capsule_service(ctx) + assert captured["model_provider"] == "custom" + assert captured["fallback_model_name"] == "vendor/fallback" + assert captured["fallback_model_provider"] == "custom" + + +def test_capsule_fallback_block_provider_wins_over_knob(): + """The summarizer fallback block's own provider wins over the chunking + 'provider' knob — fallback blocks resolve exactly like every other node's + fallback (see _resolve_endpoint(use_fallback=True)).""" + captured = {} + + class RecordingLens(StepCapsuleLens): + def __init__(self, *args, **kwargs): + super().__init__(*args, **kwargs) + captured.update( + model_provider=self._model_provider, + fallback_model_name=self._fallback_model_name, + fallback_model_provider=self._fallback_model_provider, + ) + + mgr = HistoryChunkManager( + chunking_config=SimpleNamespace( + max_steps=12, + target_source_tokens=2000, + model="tensorx/deepseek/deepseek-v4-flash", + provider="custom", + max_chunks=8, + ), + ) + ctx = SimpleNamespace( + llm_config=SimpleNamespace( + summarizer=SimpleNamespace( + fallback=SimpleNamespace(provider="google", model="gemini-2.0-flash") + ) + ) + ) + with patch.object(chunking_module, "StepCapsuleLens", RecordingLens): + mgr._build_capsule_service(ctx) + assert captured["model_provider"] == "custom" + assert captured["fallback_model_name"] == "gemini-2.0-flash" + assert captured["fallback_model_provider"] == "google" @pytest.mark.asyncio diff --git a/tests/unit/services/test_llm_lens_routing.py b/tests/unit/services/test_llm_lens_routing.py new file mode 100644 index 00000000..33f9cda0 --- /dev/null +++ b/tests/unit/services/test_llm_lens_routing.py @@ -0,0 +1,170 @@ +# Copyright 2026 Google LLC +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +"""Unit tests for provider-aware lens-model routing (get_lens_llm). + +Covers the routing matrix behind the VisualStepSummarizer 404 fix: the +Gemini-name shortcut, the explicit provider knob, summarizer-node provider +inheritance, and the diagnosability warning for non-Gemini models routed to +the Gemini API. +""" + +from unittest.mock import Mock, patch + +import pytest +from langchain_openai import ChatOpenAI + +from artemis.config.llm import LLM, LLMConfig, LLMWithFallback +from artemis.services.llm import ( + _inherit_lens_provider, + _warn_if_non_google_routes_to_google, + get_lens_llm, +) +from artemis.llm.router import ModelProvider + + +def _llm_cfg(provider: str, model: str = "m") -> LLMConfig: + """Minimal LLMConfig with only the summarizer node under test.""" + cfg = Mock(spec=LLMConfig) + cfg.summarizer = LLMWithFallback( + provider=provider, model=model, fallback=LLM(provider=provider, model="f") + ) + return cfg + + +@pytest.fixture +def capture_create(monkeypatch): + """Capture ModelFactory.create_model calls instead of building clients.""" + calls = [] + + def fake_create(endpoint): + calls.append(endpoint) + return Mock() + + monkeypatch.setattr("artemis.services.llm.ModelFactory.create_model", fake_create) + return calls + + +class TestGeminiNameShortcut: + def test_bare_gemini_name_routes_to_google(self, capture_create): + get_lens_llm(None, "gemini-2.5-flash-lite") + (ep,) = capture_create + assert ep.provider == ModelProvider.GOOGLE + assert ep.model_name == "gemini-2.5-flash-lite" + + def test_google_prefixed_name_routes_to_google(self, capture_create): + get_lens_llm(None, "google/gemini-3.8-flash") + (ep,) = capture_create + assert ep.provider == ModelProvider.GOOGLE + assert ep.model_name == "google/gemini-3.8-flash" + + def test_gemini_prefixed_name_routes_to_google(self, capture_create): + get_lens_llm(None, "gemini/gemini-2.5-flash-lite") + (ep,) = capture_create + assert ep.provider == ModelProvider.GOOGLE + + +class TestExplicitProviderKnob: + def test_explicit_custom_provider_for_namespaced_model(self, capture_create): + get_lens_llm(None, "tensorx/deepseek/deepseek-v4-flash", "custom") + (ep,) = capture_create + assert ep.provider == ModelProvider.CUSTOM + # The full model ID passes through verbatim — never prefix-parsed. + assert ep.model_name == "tensorx/deepseek/deepseek-v4-flash" + + def test_explicit_provider_beats_gemini_shortcut(self, capture_create): + # An explicit knob wins even for a Gemini name (provider-hosted + # gemini-compatible deployments). + get_lens_llm(None, "gemini-2.5-flash-lite", "openai") + (ep,) = capture_create + assert ep.provider == ModelProvider.OPENAI + + def test_unknown_provider_raises(self, capture_create): + with pytest.raises(ValueError, match="Unknown LLM provider"): + get_lens_llm(None, "some-model", "notaprovider") + + def test_temperature_and_timeout_forwarded(self, capture_create): + get_lens_llm(None, "m", "custom", temperature=0.4, timeout=12.0) + (ep,) = capture_create + assert ep.temperature == 0.4 + assert ep.timeout_seconds == 12.0 + + +class TestInheritedProvider: + def test_inherits_summarizer_node_provider(self, capture_create): + ctx = Mock() + ctx.llm_config = _llm_cfg("custom") + get_lens_llm(ctx, "tensorx/deepseek/deepseek-v4-flash") + (ep,) = capture_create + assert ep.provider == ModelProvider.CUSTOM + + def test_inherits_google_default(self, capture_create): + ctx = Mock() + ctx.llm_config = _llm_cfg("google") + get_lens_llm(ctx, "vendor/whatever") + (ep,) = capture_create + assert ep.provider == ModelProvider.GOOGLE + + def test_no_ctx_llm_config_falls_back_to_default_config(self, capture_create, monkeypatch): + ctx = Mock() + ctx.llm_config = None + monkeypatch.setattr( + "artemis.config.llm.get_default_llm_config", + lambda: _llm_cfg("openai"), + ) + get_lens_llm(ctx, "some/model") + (ep,) = capture_create + assert ep.provider == ModelProvider.OPENAI + + def test_resolution_failure_defaults_to_custom(self, capture_create, monkeypatch): + ctx = Mock() + ctx.llm_config = None + monkeypatch.setattr( + "artemis.services.llm.get_default_llm_config", + Mock(side_effect=RuntimeError("boom")), + raising=False, + ) + get_lens_llm(ctx, "some/model") + (ep,) = capture_create + assert ep.provider == ModelProvider.CUSTOM + + +class TestNonGoogleWarning: + def test_warns_on_non_gemini_name_to_google(self, caplog): + with caplog.at_level("WARNING", logger="artemis.services.llm"): + _warn_if_non_google_routes_to_google(ModelProvider.GOOGLE, "tensorx/deepseek/v4") + assert any("routed to the Google" in r.message for r in caplog.records) + + def test_no_warning_for_gemini_name(self, caplog): + with caplog.at_level("WARNING", logger="artemis.services.llm"): + _warn_if_non_google_routes_to_google(ModelProvider.GOOGLE, "gemini-2.5-flash-lite") + assert not caplog.records + + def test_no_warning_for_non_google_provider(self, caplog): + with caplog.at_level("WARNING", logger="artemis.services.llm"): + _warn_if_non_google_routes_to_google(ModelProvider.CUSTOM, "tensorx/deepseek/v4") + assert not caplog.records + + +class TestIntegration: + def test_tensorx_model_builds_chatopenai_with_gateway_base_url(self, monkeypatch): + """End-to-end: the reported 404 model now routes to the OpenAI gateway.""" + from artemis.services.llm import ModelFactory + + monkeypatch.setattr("artemis.llm.router.ModelFactory._cache", {}, raising=False) + with patch.dict("os.environ", {"OPENAI_API_KEY": "test-key"}): + llm = get_lens_llm(None, "tensorx/deepseek/deepseek-v4-flash") + assert isinstance(llm, ChatOpenAI) + assert llm.model_name == "tensorx/deepseek/deepseek-v4-flash" + assert llm.openai_api_base == "https://api.opper.ai/v3/compat/" diff --git a/tests/unit/test_memory_config.py b/tests/unit/test_memory_config.py index eb6e3579..e9c4d043 100644 --- a/tests/unit/test_memory_config.py +++ b/tests/unit/test_memory_config.py @@ -145,6 +145,35 @@ def test_legacy_default_does_not_override_memory_default(): assert cfg.memory.runtime.retry_limit == 3 +def test_lens_provider_knobs_parse_and_default_to_none(): + """The new provider knobs on step_summarizer and chunking parse from config + and default to None (auto-resolution).""" + cfg = AgentGlobalConfig.model_validate( + { + "flash": { + "step_summarizer": { + "model": "tensorx/deepseek/deepseek-v4-flash", + "provider": "custom", + } + }, + "memory": {"chunking": {"model": "vendor/some-model", "provider": "openai"}}, + } + ) + assert cfg.flash.step_summarizer.model == "tensorx/deepseek/deepseek-v4-flash" + assert cfg.flash.step_summarizer.provider == "custom" + assert cfg.memory.chunking.model == "vendor/some-model" + assert cfg.memory.chunking.provider == "openai" + + +def test_lens_provider_knob_rejects_unknown_provider(): + with pytest.raises(Exception, match="Unknown LLM provider"): + AgentGlobalConfig.model_validate( + {"flash": {"step_summarizer": {"provider": "notaprovider"}}} + ) + with pytest.raises(Exception, match="Unknown LLM provider"): + AgentGlobalConfig.model_validate({"memory": {"chunking": {"provider": "notaprovider"}}}) + + def test_recall_and_similarity_defaults(): """M4 defaults: recall block on, bounded; similarity hint on. The distance threshold was calibrated to 5 in M5 (2026-09-01 on-device dHash data: