diff --git a/apps/admin_console/routers/sessions.py b/apps/admin_console/routers/sessions.py index 6d828b77..71849792 100644 --- a/apps/admin_console/routers/sessions.py +++ b/apps/admin_console/routers/sessions.py @@ -119,9 +119,14 @@ def _list_sessions_sync(): sess_profile = model_service.resolve_session_profile( row_dict, None, state.current_profile, agent_names=agent_names ) + llm_model, llm_provider = model_service.resolve_session_llm_override(row_dict) + # A stored override is reported even when the profile stays unresolved + # (the architecture stays Flash until the trace-name pass below settles + # it), because the pinned model is a fact about the run. Legacy rows + # without the keys keep the global default. row_dict["model_info"] = ( - model_service.get_active_model_info(sess_profile) - if sess_profile + model_service.get_active_model_info(sess_profile, llm_model, llm_provider) + if sess_profile or llm_model or llm_provider else default_model_info ) if not sess_profile: @@ -143,7 +148,10 @@ def _list_sessions_sync(): row_dict, llm_traces, state.current_profile ) if sess_profile: - row_dict["model_info"] = model_service.get_active_model_info(sess_profile) + llm_model, llm_provider = model_service.resolve_session_llm_override(row_dict) + row_dict["model_info"] = model_service.get_active_model_info( + sess_profile, llm_model, llm_provider + ) if orphaned_ids: try: @@ -169,11 +177,26 @@ def _list_sessions_sync(): @router.get("/api/sessions/{session_id}") async def get_session_details(session_id: str): - """Retrieve details for a single automation session.""" + """Retrieve details for a single automation session. + + Also echoes the per-task LLM override the run recorded in device_info (the + SDK writes it there), unpacking the same JSON ``get_session_usage`` reads, so + clients see which model the session actually used. Rows written before the + override existed report ``null`` and keep the configured model. + """ row = session_repo.get_session_by_id(session_id) if not row: raise HTTPException(status_code=404, detail=f"Session {session_id} not found") - return dict(row) + payload = dict(row) + llm_model, llm_provider = model_service.resolve_session_llm_override(payload) + payload["llm_model"] = llm_model + payload["llm_provider"] = llm_provider + # model_info for parity with the list endpoint; the profile comes from the + # stored device_info, so no extra trace lookups are needed here. + payload["model_info"] = model_service.get_active_model_info( + model_service.resolve_session_profile(payload), llm_model, llm_provider + ) + return payload @router.get("/api/sessions/{session_id}/usage") diff --git a/apps/admin_console/routers/tasks.py b/apps/admin_console/routers/tasks.py index 26fca461..fcc342a2 100644 --- a/apps/admin_console/routers/tasks.py +++ b/apps/admin_console/routers/tasks.py @@ -104,6 +104,10 @@ async def run_task(request: RunRequest): task_payload.setdefault("session_id", requested_sid) task_payload.setdefault("goal", incoming_goals[0]) task_payload.setdefault("profile", request.profile or "flash") + # setdefault, not assignment: an idempotent retry replays the first + # override, so a task already queued/active keeps the LLM it started with. + task_payload.setdefault("llm_model", request.llm_model) + task_payload.setdefault("llm_provider", request.llm_provider) task_payload.setdefault("device_serial", request.device_serial) task_payload.setdefault("status", "running" if is_active else "queued") return { @@ -172,6 +176,8 @@ async def run_task(request: RunRequest): enable_outputter=request.enable_outputter, verification_level=request.verification_level, explorer_mode=request.explorer_mode, + llm_model=request.llm_model, + llm_provider=request.llm_provider, locked_app_package=request.locked_app_package, app_path=request.app_path, device_serial=target_serial, @@ -197,6 +203,16 @@ async def get_run_defaults(): } +@router.get("/api/llm-options") +async def get_llm_options(): + """Provider allowlist, ``artemis.jsonc`` presets and the configured default. + + The model picker builds its dropdowns from this instead of hard-coding a + provider list, so a new provider or preset shows up without a UI change. + """ + return model_service.get_llm_options() + + @router.get("/api/devices") async def list_devices(): """List all connected Android devices with their busy / idle status.""" @@ -305,17 +321,27 @@ async def get_status(): or state.current_profile or (running_task.get("profile") if running_task and not global_owner else None) ) - if not active_profile and (running_sid or latest_session_id): - check_sid = running_sid or latest_session_id - sess_row = session_repo.get_session_by_id(check_sid) - if sess_row: - llm_traces = session_repo.get_llm_traces_for_profile(check_sid) - agent_names = session_repo.get_agent_trace_names(check_sid) - active_profile = model_service.resolve_session_profile( - sess_row, llm_traces, agent_names=agent_names - ) - - model_info = model_service.get_active_model_info(active_profile) + # One row fetch, reused for both consumers below: the model echo needs it + # whenever a session exists, and profile resolution falls back to it when the + # worker left no profile behind. Hence the name - it feeds model resolution, + # not just the profile. No second query is issued for the override. + sess_row_for_model: dict[str, Any] | None = None + check_sid = running_sid or latest_session_id + if check_sid: + sess_row_for_model = session_repo.get_session_by_id(check_sid) + + if not active_profile and sess_row_for_model: + llm_traces = session_repo.get_llm_traces_for_profile(check_sid) + agent_names = session_repo.get_agent_trace_names(check_sid) + active_profile = model_service.resolve_session_profile( + sess_row_for_model, llm_traces, agent_names=agent_names + ) + + # The stored override applies even when active_profile came from the owner + # connection, the worker state or the queue item: a pinned model is a fact + # about the run, independent of how the profile was resolved. + llm_model, llm_provider = model_service.resolve_session_llm_override(sess_row_for_model or {}) + model_info = model_service.get_active_model_info(active_profile, llm_model, llm_provider) # Unified Global Queue: merge web tasks and external SDK/CLI device queue tickets global_queued = DeviceExecutionLock.get_queued_tasks() @@ -358,7 +384,19 @@ async def get_status(): conn_info = state.active_connections[str(latest_session_id)] is_paused = state.is_paused conn_profile = conn_info.get("profile") or active_profile - conn_model_info = model_service.get_active_model_info(conn_profile) + # This branch reports the connection's own session, so reuse the row read + # above only when it describes that same session; otherwise read it here. + # Just the one indexed lookup - the trace-name profile lookups stay on the + # main path, so this branch stays cheap. + conn_row = sess_row_for_model + if not conn_row or str(conn_row.get("session_id") or "") != str(latest_session_id): + conn_row = session_repo.get_session_by_id(latest_session_id) + conn_llm_model, conn_llm_provider = model_service.resolve_session_llm_override( + conn_row or {} + ) + conn_model_info = model_service.get_active_model_info( + conn_profile, conn_llm_model, conn_llm_provider + ) return { "status": "paused" if is_paused else "running", "paused_error": state.paused_error if is_paused else None, diff --git a/apps/admin_console/schemas/task_schema.py b/apps/admin_console/schemas/task_schema.py index 144a9720..05e353c3 100644 --- a/apps/admin_console/schemas/task_schema.py +++ b/apps/admin_console/schemas/task_schema.py @@ -12,7 +12,9 @@ # See the License for the specific language governing permissions and # limitations under the License. -from pydantic import BaseModel +from pydantic import BaseModel, model_validator + +from artemis.config.llm_override import normalize_llm_override class RunRequest(BaseModel): @@ -26,6 +28,11 @@ class RunRequest(BaseModel): # version used by the Operator ('flash' | 'pro' | 'ultra'). verification_level: str | None = None explorer_mode: str | None = None + # Per-task LLM override: pins every model of this one task to `llm_model` + # (optionally on `llm_provider`) without touching artemis.jsonc or any + # global state. Omitted/blank means "use the configured nodes". + llm_model: str | None = None + llm_provider: str | None = None locked_app_package: str | None = None app_path: str | None = None device_serial: str | None = None @@ -33,6 +40,19 @@ class RunRequest(BaseModel): session_id: str | None = None conversation_id: str | None = None + @model_validator(mode="after") + def _normalize_llm_override(self) -> "RunRequest": + """Normalise the override and reject an unusable one up front. + + Blank means "not requested", the provider is lower-cased against the + supported set, and a provider without a model is refused so FastAPI + answers 422 instead of the task dying mid-run on a broken endpoint. + """ + self.llm_model, self.llm_provider = normalize_llm_override( + self.llm_model, self.llm_provider + ) + return self + class ReplayRequest(BaseModel): device_id: str diff --git a/apps/admin_console/services/model_service.py b/apps/admin_console/services/model_service.py index 0ccc0bc0..d5ade4c8 100644 --- a/apps/admin_console/services/model_service.py +++ b/apps/admin_console/services/model_service.py @@ -15,7 +15,7 @@ import json import logging import time -from typing import Any +from typing import Any, get_args logger = logging.getLogger(__name__) @@ -55,11 +55,31 @@ def _get_llm_provider_and_model(cls) -> tuple[str, str]: return provider, model_id @classmethod - def get_active_model_info(cls, profile: str | None = None) -> dict[str, str]: - """Return active architecture and underlying LLM model configuration.""" + def get_active_model_info( + cls, + profile: str | None = None, + llm_model: str | None = None, + llm_provider: str | None = None, + ) -> dict[str, str]: + """Return active architecture and underlying LLM model configuration. + + Args: + profile: Resolved architecture ('flash' / 'pro') or None for Flash. + llm_model: Per-task model override recorded for the session. It wins + the ``id`` so the reported model is the one that actually ran. + llm_provider: Provider of that override. Only replaces the + configured provider when the task pinned it; without it each + node keeps its own provider. + """ # 1. Determine underlying LLM model and provider from config (cached) provider, model_id = cls._get_llm_provider_and_model() + # 1b. A recorded per-task override outranks the configured defaults. + if llm_model: + model_id = str(llm_model) + if llm_provider: + provider = str(llm_provider) + # 2. Determine agent architecture name (Flash vs Pro) arch_name = "Flash" if profile: @@ -78,6 +98,71 @@ def get_active_model_info(cls, profile: str | None = None) -> dict[str, str]: "architecture": f"ARTEMIS {arch_name}", } + @classmethod + def get_llm_options(cls) -> dict[str, Any]: + """Providers, the raw ``artemis.jsonc`` presets, and the configured default. + + Feeds the Console model picker: the provider allowlist is the same + ``LLMProvider`` literal the override validates against, and the presets + are read straight from the config so the choices stay in sync with + ``artemis.jsonc`` instead of being duplicated in the UI. An unreadable + config yields no presets rather than an error, matching how the rest of + the display paths degrade. + """ + from artemis.config.constants import ARTEMIS_CONFIG_FILENAME, LLMProvider + from artemis.config.paths import get_config_path + from artemis.utils.file import load_jsonc + + provider, model = cls._get_llm_provider_and_model() + presets: list[dict[str, str]] = [] + try: + with open(get_config_path(ARTEMIS_CONFIG_FILENAME), encoding="utf-8") as f: + raw = load_jsonc(f) + except Exception as exc: + logger.warning("Could not read LLM presets for display: %s", exc) + raw = None + block = raw.get("presets") if isinstance(raw, dict) else None + if isinstance(block, dict): + for name, preset in block.items(): + if isinstance(preset, dict): + presets.append( + { + "name": str(name), + "provider": str(preset.get("provider") or ""), + "model": str(preset.get("model") or ""), + } + ) + return { + "providers": list(get_args(LLMProvider)), + "presets": presets, + "default": {"provider": provider, "model": model}, + } + + @staticmethod + def resolve_session_llm_override(row_dict: dict[str, Any]) -> tuple[str | None, str | None]: + """Read the per-task LLM override a session recorded in its device_info. + + Reads the same schemaless JSON ``resolve_session_profile`` parses. Rows + written before the override existed carry neither key, which is reported + as "no override" so the caller falls back to the configured model. + """ + d_info_raw = row_dict.get("device_info") + if not d_info_raw: + return (None, None) + try: + d_info = json.loads(d_info_raw) if isinstance(d_info_raw, str) else d_info_raw + except (ValueError, TypeError): + # Malformed device_info JSON: treat the run as un-overridden. + return (None, None) + if not isinstance(d_info, dict): + return (None, None) + model = d_info.get("llm_model") + provider = d_info.get("llm_provider") + return ( + str(model).strip() or None if isinstance(model, str) else None, + str(provider).strip().lower() or None if isinstance(provider, str) else None, + ) + @staticmethod def resolve_session_profile( row_dict: dict[str, Any], diff --git a/apps/admin_console/services/task_queue_service.py b/apps/admin_console/services/task_queue_service.py index b99b4c91..e0c0bdfa 100644 --- a/apps/admin_console/services/task_queue_service.py +++ b/apps/admin_console/services/task_queue_service.py @@ -41,6 +41,7 @@ TEST_OUTPUTS_DIR, WORKSPACE_ROOT, ) +from artemis.config.llm_override import normalize_llm_override from artemis.runtime import ( AdbEndpoint, AdbTarget, @@ -505,6 +506,8 @@ def _build_worker_invocation( enable_outputter = task_item.get("enable_outputter") verification_level = task_item.get("verification_level") explorer_mode = task_item.get("explorer_mode") + llm_model = task_item.get("llm_model") + llm_provider = task_item.get("llm_provider") locked_app = task_item.get("locked_app_package") or task_item.get("locked_app") app_path = task_item.get("app_path") @@ -551,6 +554,10 @@ def _build_worker_invocation( cmd.extend(["--verification-level", str(verification_level)]) if explorer_mode: cmd.extend(["--explorer-pro-mode", str(explorer_mode)]) + if llm_model: + cmd.extend(["--model", str(llm_model)]) + if llm_provider: + cmd.extend(["--provider", str(llm_provider)]) if locked_app: cmd.extend(["--locked-app", str(locked_app)]) if app_path: @@ -1004,6 +1011,8 @@ def _create_queue_item( conversation_id: str | None, verification_level: str | None = None, explorer_mode: str | None = None, + llm_model: str | None = None, + llm_provider: str | None = None, ) -> dict[str, Any]: """Reserve a device slot and build one pending queue item for a goal.""" sess_id = single_session_id if single_session_id else str(uuid.uuid4()) @@ -1025,6 +1034,8 @@ def _create_queue_item( "enable_outputter": enable_outputter, "verification_level": verification_level, "explorer_mode": explorer_mode, + "llm_model": llm_model, + "llm_provider": llm_provider, "locked_app_package": locked_app_package, "app_path": app_path, "device_serial": assigned_serial, @@ -1052,17 +1063,28 @@ async def enqueue_tasks( conversation_id: str | None = None, verification_level: str | None = None, explorer_mode: str | None = None, + llm_model: str | None = None, + llm_provider: str | None = None, ) -> dict[str, Any]: """Enqueues one or more goals and wakes up the background worker. ``verification_level`` and ``explorer_mode`` are Pro-profile tuning knobs forwarded to the worker as ``--verification-level`` / ``--explorer-pro-mode``; they are normalised here so the queue item and the CLI see one spelling. + + ``llm_model`` / ``llm_provider`` are the per-task LLM override: they are + persisted on the queue item and forwarded to the worker as ``--model`` / + ``--provider``, which pin the models of that one task without editing + ``artemis.jsonc`` or restarting anything. A blank value means "unset". """ verification_level = ( str(verification_level).strip().lower() or None if verification_level else None ) explorer_mode = str(explorer_mode).strip().lower() or None if explorer_mode else None + # Model identifiers are case-sensitive, so only whitespace is trimmed; + # provider names are normalised to lower case and validated, so an + # unusable override is rejected before the worker is woken. + llm_model, llm_provider = normalize_llm_override(llm_model, llm_provider) cls.ensure_worker_running() enqueued_tasks = [] @@ -1105,6 +1127,8 @@ async def enqueue_tasks( conversation_id, verification_level=verification_level, explorer_mode=explorer_mode, + llm_model=llm_model, + llm_provider=llm_provider, ) state.queue_items.append(task_item) enqueued_tasks.append(task_item) diff --git a/apps/showcase_ui/src/app/components/agent-stream/agent-stream.component.html b/apps/showcase_ui/src/app/components/agent-stream/agent-stream.component.html index 93172d38..e80fb71b 100644 --- a/apps/showcase_ui/src/app/components/agent-stream/agent-stream.component.html +++ b/apps/showcase_ui/src/app/components/agent-stream/agent-stream.component.html @@ -83,6 +83,14 @@ {{ getTaskStatus(session) | uppercase }}
+ @if (getTaskModelLabel(session); as model) { + + memory + {{ model.label }} + + } @if (getDeviceSerial(session)) { phone_android @@ -151,6 +159,14 @@ {{ getTaskStatus(session) | uppercase }}
+ @if (getTaskModelLabel(session); as model) { + + memory + {{ model.label }} + + } @if (getDeviceSerial(session)) { phone_android diff --git a/apps/showcase_ui/src/app/components/agent-stream/agent-stream.component.scss b/apps/showcase_ui/src/app/components/agent-stream/agent-stream.component.scss index cb1c34f7..5432b144 100644 --- a/apps/showcase_ui/src/app/components/agent-stream/agent-stream.component.scss +++ b/apps/showcase_ui/src/app/components/agent-stream/agent-stream.component.scss @@ -3805,6 +3805,46 @@ $stream-row-height: 25px; } } + .item-model { + display: inline-flex; + align-items: center; + gap: 2.5px; + min-width: 0; + flex-shrink: 1; + font-size: 10.5px; + color: #64748b; + line-height: 1; + cursor: default; + transition: color 0.15s ease, opacity 0.15s ease; + + .model-icon { + font-size: 12px; + width: 12px; + height: 12px; + line-height: 12px; + color: #64748b; + opacity: 0.75; + flex-shrink: 0; + } + + .model-name { + overflow: hidden; + text-overflow: ellipsis; + white-space: nowrap; + min-width: 0; + font-family: inherit; + letter-spacing: -0.15px; + } + + &:hover { + color: #334155; + .model-icon { + color: #334155; + opacity: 1; + } + } + } + .item-time-divider { font-size: 10px; color: #94a3b8; diff --git a/apps/showcase_ui/src/app/components/agent-stream/agent-stream.component.ts b/apps/showcase_ui/src/app/components/agent-stream/agent-stream.component.ts index 43723721..67105d24 100644 --- a/apps/showcase_ui/src/app/components/agent-stream/agent-stream.component.ts +++ b/apps/showcase_ui/src/app/components/agent-stream/agent-stream.component.ts @@ -216,6 +216,8 @@ import { import { drawActionCoordinatesOnOverlay } from '../../utils/image-overlay.util'; +import { getTaskModelLabel as resolveTaskModelLabel, TaskModelLabel } from '../../utils/task-model-label.util'; + import { consolidateLogsToBlocks, isBackendSwitchNote, @@ -916,6 +918,18 @@ export class AgentStreamComponent implements AfterViewInit { return resolved; } + /** + * Model shown on a task card, in the "provider · model" form: the per-task + * override recorded on the queue item when there is one, otherwise the model + * persisted with the session, otherwise the globally active model. A queued + * task without an override has no session row yet and resolves the configured + * default when the worker dispatches it, which is what activeModel reports, so + * the badge shows the model the task will run with. Null when none is known. + */ + public getTaskModelLabel(session: Session): TaskModelLabel | null { + return resolveTaskModelLabel(session, this.agentService.activeModel()); + } + public selectTask(sessionId: string, event?: Event): void { if (event) event.stopPropagation(); this.agentService.selectSession(sessionId, true); diff --git a/apps/showcase_ui/src/app/components/chat-interface/chat-interface.component.html b/apps/showcase_ui/src/app/components/chat-interface/chat-interface.component.html index efd8882d..218dfc95 100644 --- a/apps/showcase_ui/src/app/components/chat-interface/chat-interface.component.html +++ b/apps/showcase_ui/src/app/components/chat-interface/chat-interface.component.html @@ -80,6 +80,14 @@ {{ getTaskStatus(session) | uppercase }}
+ @if (getTaskModelLabel(session); as model) { + + memory + {{ model.label }} + + } @if (getDeviceSerial(session)) { phone_android @@ -150,6 +158,14 @@ {{ getTaskStatus(session) | uppercase }}
+ @if (getTaskModelLabel(session); as model) { + + memory + {{ model.label }} + + } @if (getDeviceSerial(session)) { phone_android diff --git a/apps/showcase_ui/src/app/components/chat-interface/chat-interface.component.scss b/apps/showcase_ui/src/app/components/chat-interface/chat-interface.component.scss index 2de9b953..d91d7065 100644 --- a/apps/showcase_ui/src/app/components/chat-interface/chat-interface.component.scss +++ b/apps/showcase_ui/src/app/components/chat-interface/chat-interface.component.scss @@ -458,6 +458,46 @@ $text-muted: #a1a1aa; } } + .task-model { + display: inline-flex; + align-items: center; + gap: 2.5px; + min-width: 0; + flex-shrink: 1; + font-size: 11px; + color: #94a3b8; + line-height: 1; + cursor: default; + transition: color 0.15s ease, opacity 0.15s ease; + + .model-icon { + font-size: 12.5px; + width: 12.5px; + height: 12.5px; + line-height: 12.5px; + color: #94a3b8; + opacity: 0.75; + flex-shrink: 0; + } + + .model-name { + overflow: hidden; + text-overflow: ellipsis; + white-space: nowrap; + min-width: 0; + font-family: inherit; + letter-spacing: -0.15px; + } + + &:hover { + color: #64748b; + .model-icon { + color: #64748b; + opacity: 1; + } + } + } + .task-time-divider { font-size: 10px; color: #cbd5e1; diff --git a/apps/showcase_ui/src/app/components/chat-interface/chat-interface.component.ts b/apps/showcase_ui/src/app/components/chat-interface/chat-interface.component.ts index 45b36a1e..de110e22 100644 --- a/apps/showcase_ui/src/app/components/chat-interface/chat-interface.component.ts +++ b/apps/showcase_ui/src/app/components/chat-interface/chat-interface.component.ts @@ -21,6 +21,7 @@ import { AgentService } from '../../services/agent.service'; import { Session } from '../../core/models/session.model'; import { MarkdownSegment, MarkdownLine, NoteMilestone, ParsedNote } from '../../core/models/markdown.model'; import { parseNote, parseNoteLines } from '../../utils/markdown-parser.util'; +import { getTaskModelLabel as resolveTaskModelLabel, TaskModelLabel } from '../../utils/task-model-label.util'; export type { MarkdownSegment, MarkdownLine, NoteMilestone, ParsedNote }; @@ -217,6 +218,18 @@ export class ChatInterfaceComponent { return resolved; } + /** + * Model shown on a task card, in the "provider · model" form: the per-task + * override recorded on the queue item when there is one, otherwise the model + * persisted with the session, otherwise the globally active model. A queued + * task without an override has no session row yet and resolves the configured + * default when the worker dispatches it, which is what activeModel reports, so + * the badge shows the model the task will run with. Null when none is known. + */ + public getTaskModelLabel(session: Session): TaskModelLabel | null { + return resolveTaskModelLabel(session, this.agentService.activeModel()); + } + /** * Select a session in the UI to monitor its steps */ diff --git a/apps/showcase_ui/src/app/core/models/pro-tuning.model.ts b/apps/showcase_ui/src/app/core/models/pro-tuning.model.ts index 9e57c3ee..41d3cd9b 100644 --- a/apps/showcase_ui/src/app/core/models/pro-tuning.model.ts +++ b/apps/showcase_ui/src/app/core/models/pro-tuning.model.ts @@ -134,6 +134,10 @@ export const EXPLORER_MODES: readonly TuningLevel[] = [ export interface ProTuningOptions { verificationLevel?: VerificationLevelId | string; explorerMode?: ExplorerModeId | string; + /** LLM provider override for this run, sent as `llm_provider`. Unset means the server default. */ + provider?: string; + /** LLM model override for this run, sent as `llm_model`. Unset means the server default. */ + model?: string; } /** Effective defaults reported by `GET /api/run/defaults`. */ @@ -142,6 +146,30 @@ export interface ProTuningDefaults { explorer_mode?: string | null; } +/** One named model choice offered by the backend for the per-task override. */ +export interface LlmPreset { + name: string; + provider: string; + model: string; +} + +/** The model the server uses when a task requests no override. */ +export interface LlmDefault { + provider: string; + model: string; +} + +/** + * `GET /api/llm-options`: the providers and named presets the launcher can + * offer, plus the configured default. Every field is optional so a partial or + * older backend still renders the override fields. + */ +export interface LlmOptionsResponse { + providers?: string[]; + presets?: LlmPreset[]; + default?: LlmDefault | null; +} + export const DEFAULT_VERIFICATION_LEVEL: VerificationLevelId = 'final'; export const DEFAULT_EXPLORER_MODE: ExplorerModeId = 'flash'; diff --git a/apps/showcase_ui/src/app/core/models/session.model.ts b/apps/showcase_ui/src/app/core/models/session.model.ts index ef322d1c..ffc8b9de 100644 --- a/apps/showcase_ui/src/app/core/models/session.model.ts +++ b/apps/showcase_ui/src/app/core/models/session.model.ts @@ -30,6 +30,10 @@ export interface TaskQueueItem { start_time?: number; device_serial?: string | null; device_id?: string | null; + /** Per-task LLM override recorded on the queue item; null = server default. */ + llm_model?: string | null; + /** Provider of `llm_model`; null = the provider configured for the model. */ + llm_provider?: string | null; } export interface Session { @@ -44,6 +48,10 @@ export interface Session { device_serial?: string | null; device_id?: string | null; device_info?: any; + /** Per-task LLM override when the queue item carries one; null = server default. */ + llm_model?: string | null; + /** Provider of `llm_model`; null = the provider configured for the model. */ + llm_provider?: string | null; } export interface AgentStatusResponse { diff --git a/apps/showcase_ui/src/app/pages/workspace/workspace.component.html b/apps/showcase_ui/src/app/pages/workspace/workspace.component.html index a44cd84e..5316be16 100644 --- a/apps/showcase_ui/src/app/pages/workspace/workspace.component.html +++ b/apps/showcase_ui/src/app/pages/workspace/workspace.component.html @@ -167,6 +167,49 @@
+ + +
+ @if (llmOptions()) { + + } + + + @if (hasLlmOverride()) { + + } +
diff --git a/apps/showcase_ui/src/app/pages/workspace/workspace.component.scss b/apps/showcase_ui/src/app/pages/workspace/workspace.component.scss index 28535e47..2ab55e81 100644 --- a/apps/showcase_ui/src/app/pages/workspace/workspace.component.scss +++ b/apps/showcase_ui/src/app/pages/workspace/workspace.component.scss @@ -764,6 +764,110 @@ $rose-dark: #dc2626; } } } + + /* Optional LLM Provider / Model Fields, same frosted shell as the toggle */ + .llm-override-group { + display: inline-flex; + align-items: center; + gap: 4px; + + .llm-override-input { + width: 112px; + height: 22px; + box-sizing: border-box; + padding: 0 7px; + border-radius: 6px; + border: 1px solid rgba(255, 255, 255, 0.65); + background: rgba(255, 255, 255, 0.45); + backdrop-filter: blur(14px); + -webkit-backdrop-filter: blur(14px); + box-shadow: inset 0 1px 1px rgba(255, 255, 255, 0.8), 0 1px 3px rgba(0, 0, 0, 0.03); + color: #0f172a; + font-family: inherit; + font-size: 11px; + font-weight: 450; + line-height: 1; + outline: none; + transition: border-color 0.16s cubic-bezier(0.16, 1, 0.3, 1), + box-shadow 0.16s cubic-bezier(0.16, 1, 0.3, 1); + + &::placeholder { + color: #64748b; + font-size: 10.5px; + font-weight: 400; + } + + &:focus { + border-color: rgba(59, 130, 246, 0.55); + box-shadow: 0 0 0 2px rgba(59, 130, 246, 0.18), inset 0 1px 1.5px rgba(255, 255, 255, 1); + } + + &:disabled { + opacity: 0.6; + cursor: not-allowed; + } + } + + .llm-override-select { + width: 138px; + height: 22px; + box-sizing: border-box; + padding: 0 4px 0 7px; + border-radius: 6px; + border: 1px solid rgba(255, 255, 255, 0.65); + background: rgba(255, 255, 255, 0.45); + backdrop-filter: blur(14px); + -webkit-backdrop-filter: blur(14px); + box-shadow: inset 0 1px 1px rgba(255, 255, 255, 0.8), 0 1px 3px rgba(0, 0, 0, 0.03); + color: #0f172a; + font-family: inherit; + font-size: 11px; + font-weight: 450; + line-height: 1; + outline: none; + transition: border-color 0.16s cubic-bezier(0.16, 1, 0.3, 1), + box-shadow 0.16s cubic-bezier(0.16, 1, 0.3, 1); + + &:focus { + border-color: rgba(59, 130, 246, 0.55); + box-shadow: 0 0 0 2px rgba(59, 130, 246, 0.18), inset 0 1px 1.5px rgba(255, 255, 255, 1); + } + + &:disabled { + opacity: 0.6; + cursor: not-allowed; + } + } + + .llm-override-clear { + display: inline-flex; + align-items: center; + justify-content: center; + width: 18px; + height: 18px; + min-width: 18px; + padding: 0; + border: none; + border-radius: 50%; + background: rgba(255, 255, 255, 0.65); + box-shadow: inset 0 1px 1px rgba(255, 255, 255, 0.8); + color: #64748b; + cursor: pointer; + outline: none; + transition: color 0.16s cubic-bezier(0.16, 1, 0.3, 1), + background 0.16s cubic-bezier(0.16, 1, 0.3, 1); + + .material-symbols-outlined { + font-size: 12px; + line-height: 1; + } + + &:hover { + color: #0f172a; + background: rgba(255, 255, 255, 0.92); + } + } + } } } } diff --git a/apps/showcase_ui/src/app/pages/workspace/workspace.component.ts b/apps/showcase_ui/src/app/pages/workspace/workspace.component.ts index 2ecb521f..9a46afbe 100644 --- a/apps/showcase_ui/src/app/pages/workspace/workspace.component.ts +++ b/apps/showcase_ui/src/app/pages/workspace/workspace.component.ts @@ -14,13 +14,14 @@ * limitations under the License. */ -import { Component, ChangeDetectionStrategy, NgZone, DestroyRef, inject, computed, signal, ViewChild, ElementRef, OnInit } from '@angular/core'; +import { Component, ChangeDetectionStrategy, NgZone, DestroyRef, inject, computed, signal, ViewChild, ElementRef, OnInit, type WritableSignal } from '@angular/core'; import { FormsModule } from '@angular/forms'; import { AgentStreamComponent } from '../../components/agent-stream/agent-stream.component'; import { ChatInterfaceComponent } from '../../components/chat-interface/chat-interface.component'; import { FloatingVideoPlayerComponent } from '../../components/floating-video-player/floating-video-player.component'; import { AgentService } from '../../services/agent.service'; +import type { LlmOptionsResponse, LlmPreset, ProTuningOptions } from '../../core/models/pro-tuning.model'; @Component({ selector: 'app-workspace', @@ -56,6 +57,15 @@ export class WorkspaceComponent implements OnInit { public isSubmitting = signal(false); public errorMessage = signal(null); public selectedProfile = signal<'flash' | 'pro'>('flash'); + // Optional per-task LLM override, in-memory only and never restored: + // every task starts from the server default unless set here. An empty + // field is left out of the /api/run payload, so the server default applies. + public providerOverride = signal(''); + public modelOverride = signal(''); + public hasLlmOverride = computed(() => !!this.providerOverride() || !!this.modelOverride()); + // Providers/presets from GET /api/llm-options. Null when the endpoint is not + // available, which leaves the override as free-text fields only. + public llmOptions = signal(null); // Expand States (Signals for 0-latency reactivity) public isHoveringCard = signal(false); @@ -69,8 +79,19 @@ export class WorkspaceComponent implements OnInit { if (saved === 'flash' || saved === 'pro') { this.selectedProfile.set(saved); } + // The LLM override is intentionally not restored: it applies to a + // single task only. Drop keys written by earlier versions so a stale + // value can never leak into a future task. + localStorage.removeItem('artemis_provider_override'); + localStorage.removeItem('artemis_model_override'); } + // Best effort: a backend without /api/llm-options simply leaves the + // override as free-text fields. + this.agentService.getLlmOptions().subscribe((options) => { + this.llmOptions.set(options); + }); + // The global ⌘K/Ctrl+K shortcut is registered outside the Angular zone so // ordinary typing never schedules an extra change-detection pass. this.zone.runOutsideAngular(() => { @@ -95,6 +116,122 @@ export class WorkspaceComponent implements OnInit { } } + /** + * Set the optional LLM provider for the next task. Blank clears the override. + */ + public setProviderOverride(provider: string): void { + this.providerOverride.set((provider || '').trim()); + } + + /** + * Set the optional LLM model for the next task. Blank clears the override. + */ + public setModelOverride(model: string): void { + this.modelOverride.set((model || '').trim()); + } + + /** + * Drop both override fields so the next task runs on the server defaults. + */ + public clearLlmOverride(event?: MouseEvent): void { + if (event) { + event.stopPropagation(); + } + this.setProviderOverride(''); + this.setModelOverride(''); + } + + /** + * Option value for a preset from GET /api/llm-options. + */ + public llmPresetValue(index: number): string { + return `preset|${index}`; + } + + /** + * Option value for a provider from GET /api/llm-options. + */ + public llmProviderValue(provider: string): string { + return `provider|${provider}`; + } + + /** Stable track key for the preset list. */ + public trackLlmPreset(index: number, preset: LlmPreset): string { + return `${preset.provider}|${preset.model}|${index}`; + } + + /** Presets offered by the backend, empty until /api/llm-options answers. */ + public llmPresets = computed(() => this.llmOptions()?.presets || []); + + /** Providers offered by the backend, empty until /api/llm-options answers. */ + public llmProviders = computed(() => this.llmOptions()?.providers || []); + + /** + * Option matching the current override fields, empty when they are a custom + * value that is not one of the offered presets or providers. + */ + public llmSelectedOption = computed(() => { + const provider = this.providerOverride(); + const model = this.modelOverride(); + const presetIndex = this.llmPresets().findIndex( + (p) => p.model === model && (p.provider || '') === provider + ); + if (presetIndex >= 0) { + return this.llmPresetValue(presetIndex); + } + if (provider && this.llmProviders().indexOf(provider) >= 0) { + return this.llmProviderValue(provider); + } + return ''; + }); + + /** + * Text of the picker's first option: the model a task without an override + * will use while both fields are empty, or "Custom" once they are filled. + */ + public llmSelectPlaceholder = computed(() => { + if (!this.llmOptions()) { + return 'Model (optional)'; + } + if (this.hasLlmOverride()) { + return 'Custom'; + } + const fallback = this.llmOptions()?.default; + if (!fallback?.model) { + return 'Model (optional)'; + } + return fallback.provider + ? `Default: ${fallback.provider} · ${fallback.model}` + : `Default: ${fallback.model}`; + }); + + /** + * Fill the free-text override fields from the picked preset or provider. The + * fields stay editable, so anything can still be typed by hand. + */ + public applyLlmOption(value: string): void { + const options = this.llmOptions(); + if (!options || !value) { + return; + } + const separator = value.indexOf('|'); + if (separator < 0) { + return; + } + const key = value.slice(separator + 1); + if (value.startsWith('preset|')) { + const preset = this.llmPresets()[Number(key)]; + if (preset) { + this.setProviderOverride(preset.provider); + this.setModelOverride(preset.model); + } + return; + } + if (value.startsWith('provider|')) { + this.setProviderOverride(key); + } + } + /** * Computed boolean whether the currently viewed task is actively running or paused. * Only displays the stop/cancel button when inspecting an active task. @@ -154,8 +291,12 @@ export class WorkspaceComponent implements OnInit { */ public onCardClick(event: MouseEvent): void { const target = event.target as HTMLElement; - // Don't steal focus if clicking action buttons or textarea directly - if (target.closest('button') || target.tagName.toLowerCase() === 'textarea') { + // Don't steal focus if clicking action buttons, the LLM override fields + // or the textarea directly + if ( + target.closest('button, input, select') || + target.tagName.toLowerCase() === 'textarea' + ) { return; } this.focusInput(); @@ -232,24 +373,34 @@ export class WorkspaceComponent implements OnInit { } this.isInputFocused.set(false); - this.agentService.runTask(goal, this.selectedProfile()).subscribe({ - next: (res) => { - this.taskInput = ''; - if (this.dockInputRef?.nativeElement) { - this.dockInputRef.nativeElement.style.height = 'auto'; + const llmOverride: ProTuningOptions = { + provider: this.providerOverride() || undefined, + model: this.modelOverride() || undefined + }; + + this.agentService + .runTask(goal, this.selectedProfile(), undefined, undefined, llmOverride) + .subscribe({ + next: (res) => { + this.taskInput = ''; + // The override applies to this task only: a successful submit resets + // both fields so the next task falls back to the server defaults. + this.clearLlmOverride(); + if (this.dockInputRef?.nativeElement) { + this.dockInputRef.nativeElement.style.height = 'auto'; + } + this.isSubmitting.set(false); + this.agentService.fetchStatus(); + }, + error: (err) => { + console.error('Failed to submit task:', err); + this.isSubmitting.set(false); + this.errorMessage.set(err.error?.detail || 'The runner is busy. Please wait for current task to finish.'); + setTimeout(() => { + this.errorMessage.set(null); + }, 5000); } - this.isSubmitting.set(false); - this.agentService.fetchStatus(); - }, - error: (err) => { - console.error('Failed to submit task:', err); - this.isSubmitting.set(false); - this.errorMessage.set(err.error?.detail || 'The runner is busy. Please wait for current task to finish.'); - setTimeout(() => { - this.errorMessage.set(null); - }, 5000); - } - }); + }); } /** diff --git a/apps/showcase_ui/src/app/services/agent.service.ts b/apps/showcase_ui/src/app/services/agent.service.ts index b063c99b..bf61e737 100644 --- a/apps/showcase_ui/src/app/services/agent.service.ts +++ b/apps/showcase_ui/src/app/services/agent.service.ts @@ -16,10 +16,10 @@ import { Injectable, signal, inject, computed, DestroyRef, NgZone } from '@angular/core'; import { HttpClient } from '@angular/common/http'; -import { Observable } from 'rxjs'; +import { Observable, catchError, of } from 'rxjs'; import { Session, ModelInfo, TaskQueueItem, AgentStatusResponse, SessionUsage } from '../core/models/session.model'; -import { ProTuningDefaults, ProTuningOptions } from '../core/models/pro-tuning.model'; +import { LlmOptionsResponse, ProTuningDefaults, ProTuningOptions } from '../core/models/pro-tuning.model'; import { StepItemData, StepReplayFrame, LLMStreamResetEventData, StreamResetNotice, DEFAULT_STREAM_RESET_MESSAGE, PersistedCheckerStream, StreamSegment } from '../core/models/stream.model'; import { extractStepReplayFrames } from '../utils/action-formatter.util'; import { persistedStreamToSegments } from '../utils/stream-aggregator.util'; @@ -120,7 +120,10 @@ export class AgentService { ...s, status: sStatus, device_serial: serial, - model_info: isCurrentActive && this.activeModel() ? this.activeModel()! : s.model_info + // The stored model_info already reflects a per-task LLM override, so it + // wins; the global active model is only a fallback for rows that carry + // no model_info at all. + model_info: s.model_info ?? (isCurrentActive && this.activeModel() ? this.activeModel()! : s.model_info) }; sessionMap.set(s.session_id, finalSession); if (isTerminal) { @@ -407,6 +410,18 @@ export class AgentService { return this.http.get('/api/run/defaults'); } + /** + * Providers, presets and the configured default reported by + * `GET /api/llm-options`, for the launcher's per-task override picker. + * Resolves to null when the endpoint is unavailable, which leaves the + * override as free-text fields. + */ + public getLlmOptions(): Observable { + return this.http.get('/api/llm-options').pipe( + catchError(() => of(null)) + ); + } + /** Session-wide token totals, live executor context size and the run's tuning. */ public getSessionUsage(sessionId: string): Observable { return this.http.get(`/api/sessions/${encodeURIComponent(sessionId)}/usage`); @@ -442,6 +457,12 @@ export class AgentService { if (proTuning?.explorerMode) { payload.explorer_mode = proTuning.explorerMode; } + if (proTuning?.provider) { + payload.llm_provider = proTuning.provider; + } + if (proTuning?.model) { + payload.llm_model = proTuning.model; + } this.clearUserPinnedSession(); this.http.post('/api/run', payload).subscribe({ next: (res) => { @@ -1546,7 +1567,9 @@ export class AgentService { initial_goal: item.goal || '', start_time: item.start_time || item.created_at || (Date.now() / 1000 + index), status: item.status || 'pending', - device_serial: item.device_serial || item.device_id || null + device_serial: item.device_serial || item.device_id || null, + llm_model: item.llm_model || null, + llm_provider: item.llm_provider || null }; } return { diff --git a/apps/showcase_ui/src/app/utils/task-model-label.util.ts b/apps/showcase_ui/src/app/utils/task-model-label.util.ts new file mode 100644 index 00000000..6a13023d --- /dev/null +++ b/apps/showcase_ui/src/app/utils/task-model-label.util.ts @@ -0,0 +1,66 @@ +/** + * 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. + */ + +import { ModelInfo, Session } from '../core/models/session.model'; + +/** The model a task card shows, plus where that model came from. */ +export interface TaskModelLabel { + /** `provider · model`, or just whichever half is known. */ + label: string; + /** + * True when the label is what the task will resolve to rather than what it ran + * with: the card has no session row (yet), so the model comes from the + * override recorded on the queue item or from the globally active model. + */ + scheduled: boolean; +} + +/** `provider · model` from the halves that are known; null when neither is. */ +function joinModelLabel(model: string, provider: string): string | null { + if (model && provider) { + return `${provider} · ${model}`; + } + return model || provider || null; +} + +/** + * Model shown on a task card, resolved in three steps: the per-task LLM + * override recorded on the queue item, then the model persisted with the + * session, then the globally active model. A queued task without an override + * has no session row yet and resolves the configured default when the worker + * dispatches it, which is what `activeModel` reports, so its card shows the + * model the task will run with. Null when none of the three is known. + */ +export function getTaskModelLabel( + session: Session, + activeModel?: ModelInfo | null +): TaskModelLabel | null { + const storedModel = (session.model_info?.id || '').trim(); + const storedProvider = (session.model_info?.provider || '').trim(); + // The override wins; a half-filled override still pairs with the stored half + // so the label never loses half of it. + const model = (session.llm_model || '').trim() || storedModel; + const provider = (session.llm_provider || '').trim() || storedProvider; + const label = joinModelLabel(model, provider); + if (label) { + return { label, scheduled: !storedModel && !storedProvider }; + } + const activeLabel = joinModelLabel( + (activeModel?.id || '').trim(), + (activeModel?.provider || '').trim() + ); + return activeLabel ? { label: activeLabel, scheduled: true } : null; +} diff --git a/artemis/config/llm_override.py b/artemis/config/llm_override.py new file mode 100644 index 00000000..53df9453 --- /dev/null +++ b/artemis/config/llm_override.py @@ -0,0 +1,80 @@ +# 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. + +"""One normalisation/validation pass for the per-task LLM override. + +Every entry point that can carry an override (the Console API schema, the SDK +builder, the MCP tool, the MCP background worker and the queue service) +funnels the pair through +:func:`normalize_llm_override` so one typo is reported the same way everywhere, +before any trace or task exists. + +Semantics: + +* ``model`` is a raw provider model identifier and stays case-sensitive, so only + surrounding whitespace is stripped. +* ``provider`` names an endpoint, so it is stripped, lower-cased and checked + against :data:`SUPPORTED_LLM_PROVIDERS`. +* Blank/whitespace-only input means "not requested" on both sides. +* A provider without a model is rejected: a provider only says *where* to send + the model, so accepting it alone would pin nothing and fail mid-task on a + broken endpoint. +""" + +from typing import get_args + +from artemis.config.constants import LLMProvider + +#: Providers accepted by the per-task override. Single source of truth: the same +#: literal the LLM config validates against. +SUPPORTED_LLM_PROVIDERS: tuple[str, ...] = get_args(LLMProvider) + + +def normalize_llm_override( + llm_model: object = None, + llm_provider: object = None, + *, + error_prefix: str = "", +) -> tuple[str | None, str | None]: + """Normalise and validate a per-task ``(model, provider)`` override. + + Args: + llm_model: Model identifier, or anything blank/``None`` for "unset". + llm_provider: Provider name, normalised to the API's lower-case + spelling and checked against the supported set. + error_prefix: Optional prefix for the ``ValueError`` messages, used by + the SDK builder to name the method that failed. + + Returns: + The ``(model, provider)`` pair with blanks collapsed to ``None`` and + the provider lower-cased. + + Raises: + ValueError: If the provider is not a supported one, or if a provider is + given without a model. + """ + model = str(llm_model).strip() or None if llm_model is not None else None + raw_provider = str(llm_provider).strip() if llm_provider is not None else "" + provider = raw_provider.lower() or None + if provider is not None and provider not in SUPPORTED_LLM_PROVIDERS: + raise ValueError( + f"{error_prefix}unknown llm_provider {llm_provider!r}. " + "Supported providers: " + ", ".join(SUPPORTED_LLM_PROVIDERS) + ) + if provider and not model: + raise ValueError( + f"{error_prefix}llm_provider requires llm_model (got llm_provider={llm_provider!r}): " + "pass llm_model too, or drop llm_provider." + ) + return model, provider diff --git a/artemis/context.py b/artemis/context.py index fa698db0..10eb3412 100644 --- a/artemis/context.py +++ b/artemis/context.py @@ -184,6 +184,12 @@ class ArtemisContext(BaseModel): device: DeviceContext llm_config: LLMConfig | None = None + llm_model: str | None = None + """Per-task LLM override (request ``llm_model``). When set it wins over the + ``artemis.jsonc`` node config in ``artemis.services.llm._resolve_endpoint``. + Per-task by design: it never mutates the shared LLMConfig.""" + llm_provider: str | None = None + """Provider for :attr:`llm_model`; ``None`` keeps each node's provider.""" agent_config: Any = None adb_client: Any | None = None ui_adb_client: Any | None = None diff --git a/artemis/interfaces/cli/commands/batch.py b/artemis/interfaces/cli/commands/batch.py index 13f52a9f..0b7ca499 100644 --- a/artemis/interfaces/cli/commands/batch.py +++ b/artemis/interfaces/cli/commands/batch.py @@ -32,12 +32,38 @@ logger = get_logger(__name__) +async def _run_batch_goal( + agent: Agent, + goal: str, + profile_name: str, + llm_model: str | None, + llm_provider: str | None, +) -> str | dict | None: + """Runs one batch goal, pinning the LLM only when the batch asked for it. + + Without an override this stays on the plain ``agent.run_task`` call; with one + the goal goes through the request builder so the override lands on that + goal's request (and therefore only on that task) instead of the shared config. + """ + if not llm_model and not llm_provider: + return await agent.run_task(goal=goal, profile=profile_name) + request = ( + agent.new_task(goal) + .using_profile(profile_name) + .with_llm_override(model=llm_model, provider=llm_provider) + .build() + ) + return await agent.run_task(request=request) + + async def run_batch_tasks( tasks: list[str], profile_name: str = "pro", delay_seconds: float = 5.0, verification_level: str | None = None, explorer_pro_mode: str | None = None, + llm_model: str | None = None, + llm_provider: str | None = None, ) -> None: """Executes a list of automation tasks sequentially. @@ -49,6 +75,9 @@ async def run_batch_tasks( 'strict') for the Pro profile; ignored by Flash. explorer_pro_mode: Explorer tier ('flash', 'pro', 'ultra') behind ``ask_explorer`` under the Pro profile; ignored by Flash. + llm_model: Per-task model override pinned on every goal of this batch + (wins over artemis.jsonc, no restart needed). + llm_provider: Provider for ``llm_model``; requires ``llm_model``. """ if not os.environ.get("ARTEMIS_TASK_INGRESS"): os.environ["ARTEMIS_TASK_INGRESS"] = "cli" @@ -73,7 +102,7 @@ async def run_batch_tasks( status = "SUCCESS" error_msg = "" try: - result = await agent.run_task(goal=goal, profile=profile_name) + result = await _run_batch_goal(agent, goal, profile_name, llm_model, llm_provider) logger.info(f"Task {idx} completed: {result}") except Exception as e: status = "FAILED" @@ -162,6 +191,26 @@ def batch_command( help="Explorer tier behind ask_explorer under the Pro profile ('flash', 'pro', 'ultra').", ), ] = None, + llm_model: Annotated[ + str | None, + typer.Option( + "--model", + help=( + "Per-task model override applied to every LLM of every goal in the batch" + " (wins over artemis.jsonc, no restart needed), e.g. 'gemini-3.8-flash'." + ), + ), + ] = None, + llm_provider: Annotated[ + str | None, + typer.Option( + "--provider", + help=( + "Provider for --model (e.g. 'google', 'openai', 'anthropic');" + " defaults to the provider configured for each node." + ), + ), + ] = None, ) -> None: """Execute multiple automation tasks in sequence.""" task_list: list[str] = [] @@ -224,6 +273,8 @@ def batch_command( profile=profile, verification_level=verification_level, explorer_mode=explorer_pro_mode, + llm_model=llm_model, + llm_provider=llm_provider, base_url=base_url, ) if resp and resp.get("tasks"): @@ -274,5 +325,7 @@ def batch_command( delay_seconds=delay, verification_level=verification_level, explorer_pro_mode=explorer_pro_mode, + llm_model=llm_model, + llm_provider=llm_provider, ) ) diff --git a/artemis/interfaces/cli/commands/run.py b/artemis/interfaces/cli/commands/run.py index 91a5d044..33b50ec0 100644 --- a/artemis/interfaces/cli/commands/run.py +++ b/artemis/interfaces/cli/commands/run.py @@ -59,6 +59,8 @@ async def execute_task( explorer_flash_mode: str | None = None, explorer_pro_mode: str | None = None, verification_level: str | None = None, + llm_model: str | None = None, + llm_provider: str | None = None, ) -> None: """Executes a single mobile automation task end-to-end. @@ -79,6 +81,10 @@ async def execute_task( explorer_pro_mode: Override the Explorer tier for the Pro execution profile. verification_level: Coarse Checker preset ('off', 'final', 'checkpoints', 'strict'); applied before the explicit ``enable_checker`` switch. + llm_model: Per-task model override applied to every model this task + resolves; wins over ``artemis.jsonc`` nodes, needs no restart. + llm_provider: Provider for ``llm_model``; unset keeps each node's own + provider. """ effective_sid = ( session_id or os.getenv("ARTEMIS_SESSION_ID") or os.getenv("ARTEMIS_CLOUD_SESSION_ID") @@ -172,6 +178,10 @@ async def execute_task( task.with_output_description(output_description) if profile: task.using_profile(profile) + if llm_model or llm_provider: + # Scoped to this task only: the override travels on the request and + # is applied when the endpoint for each model is resolved. + task.with_llm_override(model=llm_model, provider=llm_provider) if app_path: task.with_app_path(Path(app_path)) @@ -322,6 +332,26 @@ def run_command( ), ), ] = None, + llm_model: Annotated[ + str | None, + typer.Option( + "--model", + help=( + "Per-task model override applied to every LLM of this run (wins over" + " artemis.jsonc, no restart needed), e.g. 'gemini-3.8-flash'." + ), + ), + ] = None, + llm_provider: Annotated[ + str | None, + typer.Option( + "--provider", + help=( + "Provider for --model (e.g. 'google', 'openai', 'anthropic');" + " defaults to the provider configured for each node." + ), + ), + ] = None, device_serial: Annotated[ str | None, typer.Option( @@ -385,6 +415,8 @@ def run_command( app_path=app_path, session_id=target_sid, ingress="cli", + llm_model=llm_model, + llm_provider=llm_provider, base_url=base_url, ) if resp and resp.get("tasks"): @@ -488,6 +520,8 @@ def on_status(sess_info): explorer_flash_mode=explorer_flash_mode, explorer_pro_mode=explorer_pro_mode, verification_level=verification_level, + llm_model=llm_model, + llm_provider=llm_provider, ) ) except (KeyboardInterrupt, asyncio.CancelledError): diff --git a/artemis/runtime/daemon_client.py b/artemis/runtime/daemon_client.py index 67736e65..0cd4f40e 100644 --- a/artemis/runtime/daemon_client.py +++ b/artemis/runtime/daemon_client.py @@ -221,6 +221,8 @@ def submit_task_to_daemon( conversation_id: str | None = None, verification_level: str | None = None, explorer_mode: str | None = None, + llm_model: str | None = None, + llm_provider: str | None = None, base_url: str | None = None, timeout: float = 15.0, ) -> dict[str, Any] | None: @@ -230,6 +232,10 @@ def submit_task_to_daemon( ``explorer_mode`` ('flash' | 'pro' | 'ultra') are the Pro-profile tuning knobs of ``/api/run``; they are forwarded verbatim and ignored by Flash. + ``llm_model`` / ``llm_provider`` are the optional per-task LLM override of + ``/api/run``: the Daemon pins this one task's models instead of the + ``artemis.jsonc`` nodes. They are sent as null when unset. + Returns the response JSON dict if successfully enqueued, or None on error. """ url = f"{base_url or f'http://{DEFAULT_DAEMON_HOST}:{DEFAULT_DAEMON_PORT}'}/api/run" @@ -241,6 +247,8 @@ def submit_task_to_daemon( "enable_outputter": enable_outputter, "verification_level": verification_level, "explorer_mode": explorer_mode, + "llm_model": llm_model, + "llm_provider": llm_provider, "locked_app_package": locked_app_package, "app_path": app_path, "session_id": session_id, @@ -345,12 +353,15 @@ def submit_batch_to_daemon( ingress: str = "cli", verification_level: str | None = None, explorer_mode: str | None = None, + llm_model: str | None = None, + llm_provider: str | None = None, base_url: str | None = None, timeout: float = 15.0, ) -> dict[str, Any] | None: """Submit a batch of goals to the running Daemon. - ``verification_level`` / ``explorer_mode`` apply to every goal of the batch + ``verification_level`` / ``explorer_mode`` / ``llm_model`` / + ``llm_provider`` apply to every goal of the batch (see :func:`submit_task_to_daemon`). """ url = f"{base_url or f'http://{DEFAULT_DAEMON_HOST}:{DEFAULT_DAEMON_PORT}'}/api/run" @@ -361,6 +372,8 @@ def submit_batch_to_daemon( "ingress": ingress, "verification_level": verification_level, "explorer_mode": explorer_mode, + "llm_model": llm_model, + "llm_provider": llm_provider, } try: data = json.dumps(payload).encode("utf-8") diff --git a/artemis/sdk/agent.py b/artemis/sdk/agent.py index a0028eb1..eb20dd37 100644 --- a/artemis/sdk/agent.py +++ b/artemis/sdk/agent.py @@ -562,6 +562,8 @@ async def _run_task( adb_client=self._adb_client, ui_adb_client=self._ui_adb_client, llm_config=agent_profile.llm_config, + llm_model=getattr(request, "llm_model", None), + llm_provider=getattr(request, "llm_provider", None), agent_config=self._config, ) @@ -1143,6 +1145,15 @@ def _prepare_tracing(self, task: Task, context: ArtemisContext): run_tuning = run_tuning_summary(self._config, task.request.profile) if run_tuning: device_data["run_tuning"] = run_tuning + # Echo the per-task LLM override so the console can report which model + # this run really used. device_info is a schemaless JSON dict, so there + # is nothing to migrate: rows written before the override existed simply + # lack these keys and fall back to the configured nodes. A provider can + # never appear without a model (the builder rejects that combination). + if task.request.llm_model: + device_data["llm_model"] = str(task.request.llm_model).strip() + if task.request.llm_provider: + device_data["llm_provider"] = str(task.request.llm_provider).strip().lower() target_sid = ( self._session_id diff --git a/artemis/sdk/builders/task_request_builder.py b/artemis/sdk/builders/task_request_builder.py index 706802c5..876d2f01 100644 --- a/artemis/sdk/builders/task_request_builder.py +++ b/artemis/sdk/builders/task_request_builder.py @@ -25,6 +25,7 @@ except ImportError: from typing import Self +from artemis.config.llm_override import normalize_llm_override from artemis.constants import RECURSION_LIMIT from artemis.sdk.types.agent import AgentProfile from artemis.sdk.types.task import TaskRequest, TaskRequestCommon @@ -169,6 +170,8 @@ def __init__(self, goal: str): self._name: str | None = None self._output_description = None self._output_format: type[TIn] | None = None + self._llm_model: str | None = None + self._llm_provider: str | None = None @classmethod def from_common(cls, goal: str, common: TaskRequestCommon): @@ -190,6 +193,32 @@ def using_profile(self, profile: str | AgentProfile) -> "TaskRequestBuilder[TIn] self._profile = profile return self + def with_llm_override( + self, + model: str | None = None, + provider: str | None = None, + ) -> "TaskRequestBuilder[TIn]": + """Pin the LLM of this task only, without changing global config. + + The override wins over the ``artemis.jsonc`` node configuration for + every model this task resolves, so a single run can use a different + model or provider without a restart. Omitted/blank values keep the + configured value. The override also pins the resolved fallback model. + + Args: + model: Model identifier override (e.g. ``gemini-3.8-flash``) + provider: Provider override (e.g. ``openai``, ``anthropic``) + + Raises: + ValueError: If ``provider`` is set without ``model``, or names an + unknown provider. Both are caught here instead of failing + mid-task with a broken endpoint. + """ + self._llm_model, self._llm_provider = normalize_llm_override( + model, provider, error_prefix="with_llm_override: " + ) + return self + def with_name(self, name: str) -> "TaskRequestBuilder[TIn]": """Set the name of the task - useful when recording traces. @@ -255,5 +284,7 @@ def build(self) -> TaskRequest[TIn]: llm_output_path=self._llm_output_path, locked_app_package=self._locked_app_package, app_path=self._app_path, + llm_model=self._llm_model, + llm_provider=self._llm_provider, ) return task_request diff --git a/artemis/sdk/types/task.py b/artemis/sdk/types/task.py index 7c4a0be6..b716dca0 100644 --- a/artemis/sdk/types/task.py +++ b/artemis/sdk/types/task.py @@ -133,6 +133,10 @@ class TaskRequest(TaskRequestCommon, Generic[TOutput]): execution (default: False) trace_path: Directory path to save trace data if recording is enabled llm_output_path: Path to save LLM output data + llm_model: Optional per-task LLM override pinning every model of this + task to a specific identifier (takes precedence over the configured + ``artemis.jsonc`` nodes; ``None`` keeps the configured nodes) + llm_provider: Optional provider for ``llm_model`` """ model_config = {"ignored_types": (CyFunctionDetector,)} @@ -142,6 +146,8 @@ class TaskRequest(TaskRequestCommon, Generic[TOutput]): output_description: str | None = None output_format: type[TOutput] | None = None enable_remote_tracing: bool = False + llm_model: str | None = None + llm_provider: str | None = None class TaskResult(BaseModel): diff --git a/artemis/services/llm.py b/artemis/services/llm.py index 784e971f..dcfcbcfc 100644 --- a/artemis/services/llm.py +++ b/artemis/services/llm.py @@ -1037,7 +1037,15 @@ def _resolve_endpoint( is_utils: bool = False, use_fallback: bool = False, ) -> ModelEndpoint: - """Cleanly resolves a ModelEndpoint from context llm_config.""" + """Cleanly resolves a ModelEndpoint from context llm_config. + + Precedence is the per-task override (``ctx.llm_model`` / ``ctx.llm_provider``) + over the ``artemis.jsonc`` node config over the built-in defaults. The + override is applied *after* the ``use_fallback`` unwrap on purpose: it pins + the fallback model as well, so a task that asked for a model never silently + drops to a different one mid-run. The config object itself is only read, so + other tasks sharing it keep their configured models. + """ if getattr(ctx, "llm_config", None) is None: try: ctx.llm_config = get_default_llm_config() @@ -1061,8 +1069,15 @@ def _get_val(obj, attr, expected_type): val = getattr(obj, attr, None) return val if isinstance(val, expected_type) else None - provider_val = getattr(cfg, "provider", "google") - model_val = getattr(cfg, "model", "gemini-2.5-flash") + # Precedence: per-task override > artemis.jsonc node config > built-in + # default. The override lives on the per-task context (see sdk/agent.py), so + # a task can pin a different model/provider without mutating the shared + # LLMConfig or restarting the host. + override_provider = _get_val(ctx, "llm_provider", str) + override_model = _get_val(ctx, "llm_model", str) + + provider_val = (override_provider or "").strip() or getattr(cfg, "provider", "google") + model_val = (override_model or "").strip() or getattr(cfg, "model", "gemini-2.5-flash") return ModelEndpoint( provider=ModelProvider.from_string(provider_val), diff --git a/mcp_server/background/task_runner.py b/mcp_server/background/task_runner.py index a6f377ce..85aebc44 100644 --- a/mcp_server/background/task_runner.py +++ b/mcp_server/background/task_runner.py @@ -41,6 +41,7 @@ except Exception: load_dotenv(os.path.join(PROJECT_ROOT, ".env")) +from artemis.config.llm_override import normalize_llm_override from artemis.runtime import trace_store from mcp_server.notifiers import notify from mcp_server.utils import device_utils @@ -121,6 +122,8 @@ async def run_task( device_serial: str | None = None, verification_level: str | None = None, explorer_pro_mode: str | None = None, + llm_model: str | None = None, + llm_provider: str | None = None, ): """Executes the mobile automation agent task and logs all actions/results. @@ -128,6 +131,11 @@ async def run_task( ``explorer_pro_mode`` ('flash' | 'pro' | 'ultra') are Pro-profile tuning knobs mirroring ``artemis run --verification-level / --explorer-pro-mode``; the Flash profile ignores them. + + ``llm_model`` / ``llm_provider`` are the per-task LLM override mirroring + ``artemis run --model / --provider``: they pin every model this task + resolves without touching ``artemis.jsonc``. Note ``model`` (the Flash/Pro + profile) is a different argument. """ trace_dir = trace_store.get_trace_dir(trace_id) os.makedirs(trace_dir, exist_ok=True) @@ -258,6 +266,9 @@ async def run_task( task_builder.with_output_description(description=expected_output_desc) if model.lower() == "flash": task_builder.using_profile("flash") + llm_model, llm_provider = normalize_llm_override(llm_model, llm_provider) + if llm_model or llm_provider: + task_builder.with_llm_override(model=llm_model, provider=llm_provider) result = await agent.run_task(request=task_builder.build()) print(f"Task completed. Result: {result}") @@ -434,6 +445,17 @@ async def run_task( "--explorer-pro-mode", help="Pro-profile Explorer perception version: 'flash', 'pro' or 'ultra'", ) + parser.add_argument( + "--llm-model", + help=( + "Per-task LLM model override applied to every model of this run" + " (distinct from --model, which selects the Flash/Pro profile)" + ), + ) + parser.add_argument( + "--llm-provider", + help="Provider for --llm-model ('google', 'openai', 'anthropic', ...)", + ) args = parser.parse_args() @@ -449,5 +471,7 @@ async def run_task( device_serial=args.device_serial, verification_level=args.verification_level, explorer_pro_mode=args.explorer_pro_mode, + llm_model=args.llm_model, + llm_provider=args.llm_provider, ) ) diff --git a/mcp_server/tools/task_runner.py b/mcp_server/tools/task_runner.py index d30c42eb..074b974e 100644 --- a/mcp_server/tools/task_runner.py +++ b/mcp_server/tools/task_runner.py @@ -27,6 +27,7 @@ from mcp_server.notifiers import notify from mcp_server.utils import env_utils from artemis.config import ExplorerVersion, checker_overrides_for_level +from artemis.config.llm_override import normalize_llm_override from artemis.config.runtime import read_ipc_port from artemis.runtime import ( DeviceExecutionLock, @@ -212,6 +213,8 @@ def mobile_run_task( device_serial: str | None = None, verification_level: str | None = None, explorer_mode: str | None = None, + llm_model: str | None = None, + llm_provider: str | None = None, ) -> dict[str, Any]: """Starts an autonomous mobile UI automation subagent on a connected Android device. @@ -277,6 +280,19 @@ def mobile_run_task( by the Operator: `"flash"` (1-shot detection, the default), `"pro"` (3-turn ReAct), `"ultra"` (deep pixel reasoning; slowest). Ignored for Flash. + llm_model: Optional. Per-task LLM model override applied to every + model this task resolves (e.g. `"gemini-3.8-flash"`), overriding the + `artemis.jsonc` node configuration for this run only — no restart and + no config edit. It intentionally pins the resolved fallback model as + well, so the whole task runs on the override. Leave unset to keep the + configured models. Note this is *not* the `model` profile argument + above: that one selects the Flash/Pro execution architecture, this one + selects the LLM itself, and the two combine freely. + llm_provider: Optional. Provider for `llm_model` (`"google"`, + `"openai"`, `"anthropic"`, `"openrouter"`, `"xai"`, `"vertexai"`, + `"ollama"`, `"vllm"`, `"custom"`). When omitted, each node keeps its + configured provider. Requires `llm_model`; an unknown provider or a + provider without a model is rejected before the task starts. """ # 0. Validate and normalize model if model.lower() not in ("flash", "pro"): @@ -285,6 +301,9 @@ def mobile_run_task( # 0b. Validate the Pro tuning knobs before any trace exists so a typo is a # plain tool error rather than a failed trace on disk. verification_level, explorer_mode = _normalize_pro_tuning(verification_level, explorer_mode) + # 0c. Same fail-fast treatment for the per-task LLM override, before any + # trace exists so an unusable override is a plain tool error. + llm_model, llm_provider = normalize_llm_override(llm_model, llm_provider) # 1. Generate a unique trace_id trace_id = str(uuid.uuid4()) @@ -325,6 +344,8 @@ def mobile_run_task( conversation_id=conversation_id, verification_level=verification_level, explorer_mode=explorer_mode, + llm_model=llm_model, + llm_provider=llm_provider, base_url=base_url, ) if resp and resp.get("status") == "rejected": @@ -471,6 +492,12 @@ def mobile_run_task( cmd.extend(["--verification-level", verification_level]) if explorer_mode: cmd.extend(["--explorer-pro-mode", explorer_mode]) + # Distinct from this runner's own ``--model`` flag, which carries the + # Flash/Pro profile; these carry the LLM override. + if llm_model: + cmd.extend(["--llm-model", llm_model]) + if llm_provider: + cmd.extend(["--llm-provider", llm_provider]) env = os.environ.copy() env["ARTEMIS_SESSION_ID"] = trace_id diff --git a/packages/artemis-client/src/artemis_client/client.py b/packages/artemis-client/src/artemis_client/client.py index d1d3e2d4..4ad04208 100644 --- a/packages/artemis-client/src/artemis_client/client.py +++ b/packages/artemis-client/src/artemis_client/client.py @@ -52,6 +52,27 @@ def _normalize_choice(value: str | None, name: str, choices: tuple[str, ...]) -> return normalized +def _normalize_llm_override( + llm_model: str | None, llm_provider: str | None +) -> tuple[str | None, str | None]: + """Normalise the per-task ``(model, provider)`` override before submitting. + + Blank means "not requested" rather than an override with ``""``, and the + provider is lower-cased to the API's spelling. The model identifier stays + case-sensitive, so only whitespace is trimmed. Unsupported providers are + *not* rejected here: the host validates them and answers 4xx, which keeps + this package dependency-free and forwards the server's own message. + + Mirrors ``artemis.config.llm_override.normalize_llm_override``, which cannot + be imported because this package must stay runtime-dependency-free. + Non-string input is stringified like the canonical helper so the mirror + is exact; unsupported providers are *not* rejected here. + """ + model = str(llm_model).strip() or None if llm_model is not None else None + provider = str(llm_provider).strip().lower() or None if llm_provider is not None else None + return model, provider + + class ArtemisClient: """Thin client for an Artemis daemon running on another host. @@ -186,6 +207,8 @@ async def submit( task_id: str | None = None, verification_level: VerificationLevel | None = None, explorer_mode: ExplorerMode | None = None, + llm_model: str | None = None, + llm_provider: str | None = None, options: Mapping[str, Any] | None = None, ) -> TaskHandle: """Submit one task and return immediately after scheduler admission. @@ -196,7 +219,11 @@ async def submit( ``strict``: how much the Checker audits a Pro run) and ``explorer_mode`` (``flash`` | ``pro`` | ``ultra``: the Pro Operator's perception depth) are Pro-only tuning knobs; the Flash profile ignores - them. Experimental, forward-compatible fields belong in ``options``. + them. ``llm_model`` / ``llm_provider`` are the per-task LLM override: + the host pins this task's models instead of its configured ones, with + no restart and no config change. Both are per-call (blank falls back to + "not requested") and the override also pins the fallback model. + Experimental, forward-compatible fields belong in ``options``. """ normalized_goal = goal.strip() if not normalized_goal: @@ -214,6 +241,7 @@ async def submit( raise ValueError("task_id must be a valid UUID string") from exc resolved_profile = profile or self.default_profile resolved_device = device_serial or self.device_serial + resolved_model, resolved_provider = _normalize_llm_override(llm_model, llm_provider) payload: dict[str, Any] = { "goal": normalized_goal, "profile": resolved_profile, @@ -229,6 +257,8 @@ async def submit( "conversation_id": conversation_id, "verification_level": resolved_level, "explorer_mode": resolved_mode, + "llm_model": resolved_model, + "llm_provider": resolved_provider, "options": dict(options) if options is not None else None, } payload.update({key: value for key, value in optional_values.items() if value is not None}) @@ -306,6 +336,8 @@ async def run( task_id: str | None = None, verification_level: VerificationLevel | None = None, explorer_mode: ExplorerMode | None = None, + llm_model: str | None = None, + llm_provider: str | None = None, options: Mapping[str, Any] | None = None, timeout: float = 1800.0, poll_interval: float | None = None, @@ -323,6 +355,8 @@ async def run( task_id=task_id, verification_level=verification_level, explorer_mode=explorer_mode, + llm_model=llm_model, + llm_provider=llm_provider, options=options, ) return await self.wait_for_task( @@ -341,6 +375,8 @@ async def run_task(self, task: Any, **overrides: Any) -> TaskResult: "device_serial": getattr(task, "device_serial", None) or getattr(task, "device_id", None), "locked_app_package": getattr(task, "locked_package", None), + "llm_model": getattr(task, "llm_model", None), + "llm_provider": getattr(task, "llm_provider", None), } values.update(overrides) return await self.run(goal, **values) diff --git a/packages/artemis-client/src/artemis_client/models.py b/packages/artemis-client/src/artemis_client/models.py index e9c4e2d4..dfbad6fd 100644 --- a/packages/artemis-client/src/artemis_client/models.py +++ b/packages/artemis-client/src/artemis_client/models.py @@ -89,6 +89,8 @@ class TaskResult: status: str goal: str | None = None profile: str | None = None + llm_model: str | None = None + llm_provider: str | None = None device_serial: str | None = None output: Any = None error: str | None = None @@ -141,6 +143,8 @@ def from_payload( status=(_string(payload.get("status")) or "unknown").lower(), goal=_string(payload.get("goal") or payload.get("initial_goal")), profile=_string(payload.get("profile")), + llm_model=_string(payload.get("llm_model")), + llm_provider=_string(payload.get("llm_provider")), device_serial=_device_from_payload(payload), output=output, error=_string(payload.get("error") or payload.get("error_message")), diff --git a/packages/artemis-client/tests/test_client.py b/packages/artemis-client/tests/test_client.py index a5480a7a..6328927a 100644 --- a/packages/artemis-client/tests/test_client.py +++ b/packages/artemis-client/tests/test_client.py @@ -11,6 +11,7 @@ import unittest from collections import defaultdict from collections.abc import Mapping +from types import SimpleNamespace from typing import Any from artemis_client import ( @@ -131,6 +132,154 @@ async def test_submit_rejects_unknown_pro_tuning_values_before_any_request(self) await self.client.submit("Open Settings", explorer_mode="turbo") self.assertEqual(self.transport.calls, []) + async def test_submit_forwards_per_task_llm_override(self) -> None: + task_id = "00000000-0000-4000-8000-000000000331" + self.transport.add( + "POST", + "/api/run", + {"status": "started", "tasks": [{"session_id": task_id, "status": "pending"}]}, + ) + + await self.client.submit( + "Audit checkout", + profile="pro", + task_id=task_id, + llm_model=" gemini-3.8-pro ", + llm_provider="OpenAI", + ) + + body = self.transport.calls[0][2] + assert body is not None + # Model identifiers keep their spelling; only whitespace is trimmed and + # the provider is normalised to the API's lower-case spelling. + self.assertEqual(body["llm_model"], "gemini-3.8-pro") + self.assertEqual(body["llm_provider"], "openai") + + async def test_submit_omits_llm_override_when_unset(self) -> None: + task_id = "00000000-0000-4000-8000-000000000332" + self.transport.add( + "POST", + "/api/run", + {"status": "started", "tasks": [{"session_id": task_id, "status": "pending"}]}, + ) + + await self.client.submit("Open Settings", task_id=task_id, llm_model=" ") + + body = self.transport.calls[0][2] + assert body is not None + self.assertNotIn("llm_model", body) + self.assertNotIn("llm_provider", body) + + async def test_submit_llm_override_is_per_call_only(self) -> None: + """The override never sticks to the client: it is a per-call argument.""" + transport = FakeTransport() + client = ArtemisClient("https://artemis.example.test", transport=transport) + for index in (341, 342): + transport.add( + "POST", + "/api/run", + { + "status": "started", + "tasks": [ + {"session_id": f"00000000-0000-4000-8000-{index:012d}", "status": "pending"} + ], + }, + ) + + await client.submit("Open Settings") + await client.submit("Open Settings", llm_model="gpt-5.1", llm_provider="openai") + + first, second = transport.calls[0][2], transport.calls[1][2] + assert first is not None and second is not None + self.assertNotIn("llm_model", first) + self.assertNotIn("llm_provider", first) + self.assertEqual(second["llm_model"], "gpt-5.1") + self.assertEqual(second["llm_provider"], "openai") + self.assertFalse(hasattr(client, "default_llm_model")) + self.assertFalse(hasattr(client, "default_llm_provider")) + + async def test_run_task_forwards_llm_override_from_task_and_overrides(self) -> None: + task_id = "00000000-0000-4000-8000-000000000361" + self.transport.add( + "POST", + "/api/run", + {"status": "started", "tasks": [{"session_id": task_id, "status": "pending"}]}, + ) + self.transport.add( + "GET", + f"/api/sessions/{task_id}", + { + "session_id": task_id, + "status": "completed", + "llm_model": "gpt-5.1", + "llm_provider": "openai", + }, + ) + task = SimpleNamespace( + goal="Audit checkout", + profile="pro", + llm_model="gemini-3.8-pro", + llm_provider="google", + ) + + result = await self.client.run_task(task, llm_model="gpt-5.1", llm_provider="openai") + + body = self.transport.calls[0][2] + assert body is not None + # Explicit overrides win over the values carried by the task object. + self.assertEqual(body["profile"], "pro") + self.assertEqual(body["llm_model"], "gpt-5.1") + self.assertEqual(body["llm_provider"], "openai") + # Both halves of the override are echoed back by GET /api/sessions/{id}, + # so the caller can see which model AND provider the host actually ran. + self.assertEqual(result.llm_model, "gpt-5.1") + self.assertEqual(result.llm_provider, "openai") + + async def test_run_task_uses_task_llm_override_when_not_overridden(self) -> None: + task_id = "00000000-0000-4000-8000-000000000362" + self.transport.add( + "POST", + "/api/run", + {"status": "started", "tasks": [{"session_id": task_id, "status": "pending"}]}, + ) + self.transport.add( + "GET", + f"/api/sessions/{task_id}", + {"session_id": task_id, "status": "completed"}, + ) + task = SimpleNamespace( + goal="Audit checkout", llm_model="gemini-3.8-pro", llm_provider="Google" + ) + + await self.client.run_task(task) + + body = self.transport.calls[0][2] + assert body is not None + self.assertEqual(body["llm_model"], "gemini-3.8-pro") + self.assertEqual(body["llm_provider"], "google") + + async def test_run_forwards_per_task_llm_override(self) -> None: + task_id = "00000000-0000-4000-8000-000000000351" + self.transport.add( + "POST", + "/api/run", + {"status": "started", "tasks": [{"session_id": task_id, "status": "pending"}]}, + ) + self.transport.add( + "GET", + f"/api/sessions/{task_id}", + {"session_id": task_id, "status": "completed", "llm_model": "gpt-5.1"}, + ) + + result = await self.client.run("Audit checkout", llm_model="gpt-5.1", llm_provider="openai") + + body = self.transport.calls[0][2] + assert body is not None + self.assertEqual(body["llm_model"], "gpt-5.1") + self.assertEqual(body["llm_provider"], "openai") + # The host echoes the effective model back on the task payload. + self.assertEqual(result.llm_model, "gpt-5.1") + async def test_submit_rejected_task_raises_specific_error(self) -> None: self.transport.add( "POST", diff --git a/playground/backend_manager/app/auth/jwt_handler.py b/playground/backend_manager/app/auth/jwt_handler.py index f87991fa..a23e7dee 100644 --- a/playground/backend_manager/app/auth/jwt_handler.py +++ b/playground/backend_manager/app/auth/jwt_handler.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -from datetime import datetime, timedelta, timezone +from datetime import datetime, timedelta, timezone, UTC from fastapi import Depends, HTTPException, Security, status from fastapi.security import HTTPAuthorizationCredentials, HTTPBearer from jose import JWTError, jwt @@ -23,13 +23,11 @@ def create_access_token(user_id: str, extra_claims: dict | None = None) -> str: """Generate a signed JWT access token.""" - expire = datetime.now(timezone.utc) + timedelta( - minutes=settings.JWT_ACCESS_TOKEN_EXPIRE_MINUTES - ) + expire = datetime.now(UTC) + timedelta(minutes=settings.JWT_ACCESS_TOKEN_EXPIRE_MINUTES) to_encode = { "sub": user_id, "exp": expire, - "iat": datetime.now(timezone.utc), + "iat": datetime.now(UTC), "iss": "artemis-backend-manager", } if extra_claims: diff --git a/playground/backend_manager/app/auth/otp_service.py b/playground/backend_manager/app/auth/otp_service.py index bb1f0d62..750dfdcb 100644 --- a/playground/backend_manager/app/auth/otp_service.py +++ b/playground/backend_manager/app/auth/otp_service.py @@ -12,7 +12,7 @@ # See the License for the specific language governing permissions and # limitations under the License. -from datetime import datetime, timedelta, timezone +from datetime import datetime, timedelta, timezone, UTC import logging import secrets from app.config import settings @@ -38,7 +38,7 @@ def generate_otp(self, identifier: str) -> str: else: code = f"{secrets.randbelow(900000) + 100000}" - expires_at = datetime.now(timezone.utc) + timedelta(seconds=settings.OTP_EXPIRE_SECONDS) + expires_at = datetime.now(UTC) + timedelta(seconds=settings.OTP_EXPIRE_SECONDS) self._store[ident] = { "code": code, "expires_at": expires_at, @@ -65,7 +65,7 @@ def verify_otp(self, identifier: str, code: str) -> bool: logger.warning(f"[OTP Service] No active OTP found for {ident}") return False - if datetime.now(timezone.utc) > record["expires_at"]: + if datetime.now(UTC) > record["expires_at"]: logger.warning(f"[OTP Service] OTP for {ident} has expired") self._store.pop(ident, None) return False diff --git a/playground/backend_manager/app/services/bigquery_service.py b/playground/backend_manager/app/services/bigquery_service.py index 8588ec16..8e8f64a1 100644 --- a/playground/backend_manager/app/services/bigquery_service.py +++ b/playground/backend_manager/app/services/bigquery_service.py @@ -34,7 +34,7 @@ class BigQueryMappingService: def __init__(self): self._local_cache: dict[str, SessionRecord] = {} - self._client: Optional["bigquery.Client"] = None + self._client: bigquery.Client | None = None self._table_ref: str = ( f"{settings.GCP_PROJECT_ID}.{settings.BQ_DATASET}.{settings.BQ_TABLE}" ) diff --git a/playground/backend_manager/app/services/docker_service.py b/playground/backend_manager/app/services/docker_service.py index 85d040e5..11512153 100644 --- a/playground/backend_manager/app/services/docker_service.py +++ b/playground/backend_manager/app/services/docker_service.py @@ -33,7 +33,7 @@ class DockerManagerService: """Manages creation, monitoring, and deletion of ephemeral Artemis session containers via Docker Socket.""" def __init__(self): - self._client: Optional["docker.DockerClient"] = None + self._client: docker.DockerClient | None = None if DOCKER_AVAILABLE: try: self._client = docker.DockerClient(base_url=settings.DOCKER_SOCKET_PATH) diff --git a/playground/backend_manager/app/services/session_manager.py b/playground/backend_manager/app/services/session_manager.py index 17bf6d64..c95ef296 100644 --- a/playground/backend_manager/app/services/session_manager.py +++ b/playground/backend_manager/app/services/session_manager.py @@ -13,7 +13,7 @@ # limitations under the License. import asyncio -from datetime import datetime, timedelta, timezone +from datetime import datetime, timedelta, timezone, UTC import logging import uuid from app.config import settings @@ -49,7 +49,7 @@ async def _reaper_loop(self): while True: try: await asyncio.sleep(settings.REAPER_CHECK_INTERVAL_SECONDS) - now = datetime.now(timezone.utc) + now = datetime.now(UTC) active_sessions = await bigquery_service.list_all_active_sessions() for s in active_sessions: @@ -78,7 +78,7 @@ async def _reaper_loop(self): async def create_session(self, user_id: str, request: CreateSessionRequest) -> SessionResponse: """Provision a new Cuttlefish emulator + Artemis container and link them via ADB.""" session_id = str(uuid.uuid4()) - now = datetime.now(timezone.utc) + now = datetime.now(UTC) expires_at = now + timedelta(minutes=request.ttl_minutes) logger.info( @@ -164,7 +164,7 @@ async def heartbeat(self, session_id: str, user_id: str) -> HeartbeatResponse | if not record or record.user_id != user_id or record.status == SessionStatus.TERMINATED: return None - now = datetime.now(timezone.utc) + now = datetime.now(UTC) record.last_heartbeat_at = now record.expires_at = now + timedelta(minutes=settings.SESSION_TTL_MINUTES) await bigquery_service.save_session(record) diff --git a/tests/unit/admin_console/test_model_service.py b/tests/unit/admin_console/test_model_service.py index a0634693..81ebb640 100644 --- a/tests/unit/admin_console/test_model_service.py +++ b/tests/unit/admin_console/test_model_service.py @@ -13,10 +13,12 @@ # limitations under the License. import json +from typing import get_args from unittest.mock import patch, MagicMock import pytest +from apps.admin_console.routers import tasks as tasks_router from apps.admin_console.services.model_service import ModelService @@ -51,3 +53,159 @@ def test_resolve_session_profile_from_agent_names(): row = {"device_info": None} assert ModelService.resolve_session_profile(row, agent_names=["planner", "operator"]) == "pro" assert ModelService.resolve_session_profile(row, agent_names=["flashrunner"]) == "flash" + + +# The configured (global) LLM is patched so these tests never read artemis.jsonc. +_CONFIGURED = ("google", "gemini-3.8-flash") + + +@patch.object(ModelService, "_get_llm_provider_and_model", return_value=_CONFIGURED) +def test_get_active_model_info_prefers_the_recorded_override(_configured): + """A per-task override wins the model id and, when given, the provider.""" + info = ModelService.get_active_model_info("pro", "gpt-5.1", "openai") + + assert info["id"] == "gpt-5.1" + assert info["provider"] == "openai" + assert info["name"] == "Pro" + assert info["architecture"] == "ARTEMIS Pro" + + +@patch.object(ModelService, "_get_llm_provider_and_model", return_value=_CONFIGURED) +def test_get_active_model_info_falls_back_to_config_without_override(_configured): + """Rows without the override keys keep reporting the configured model.""" + info = ModelService.get_active_model_info("flash") + + assert info["id"] == "gemini-3.8-flash" + assert info["provider"] == "google" + assert info["name"] == "Flash" + + +@patch.object(ModelService, "_get_llm_provider_and_model", return_value=_CONFIGURED) +def test_get_active_model_info_keeps_node_provider_when_override_has_none(_configured): + """A model-only override leaves the configured provider in place.""" + info = ModelService.get_active_model_info("pro", "gemini-3.8-pro") + + assert info["id"] == "gemini-3.8-pro" + assert info["provider"] == "google" + + +def test_resolve_session_llm_override_reads_device_info(): + row = {"device_info": json.dumps({"llm_model": " gpt-5.1 ", "llm_provider": " OpenAI "})} + assert ModelService.resolve_session_llm_override(row) == ("gpt-5.1", "openai") + + +def test_resolve_session_llm_override_accepts_a_parsed_device_info(): + row = {"device_info": {"llm_model": "gpt-5.1", "llm_provider": "openai"}} + assert ModelService.resolve_session_llm_override(row) == ("gpt-5.1", "openai") + + +def test_resolve_session_llm_override_is_empty_for_legacy_and_broken_rows(): + # No device_info at all, a pre-override row, and malformed JSON. + assert ModelService.resolve_session_llm_override({}) == (None, None) + assert ModelService.resolve_session_llm_override({"device_info": None}) == (None, None) + assert ModelService.resolve_session_llm_override( + {"device_info": json.dumps({"profile": "pro"})} + ) == ( + None, + None, + ) + assert ModelService.resolve_session_llm_override({"device_info": "{not json"}) == (None, None) + assert ModelService.resolve_session_llm_override({"device_info": '"a string"'}) == (None, None) + + +def test_resolve_session_llm_override_treats_blank_values_as_unset(): + row = {"device_info": json.dumps({"llm_model": " ", "llm_provider": ""})} + assert ModelService.resolve_session_llm_override(row) == (None, None) + + +# --- GET /api/llm-options: providers, raw artemis.jsonc presets, default ----- + +_PRESETS_JSONC = """ +// A comment, so the reader has to strip it. +{ + "default": { "provider": "anthropic", "model": "claude-x" }, + "presets": { + "gemini-flash": { "provider": "google", "model": "gemini-2.5-flash", + "fallback": { "provider": "google", "model": "gemini-2.0-flash" } }, + "local-ollama": { "provider": "openai", "model": "llama3.2-vision" } + } +} +""" + + +def _jsonc_config(tmp_path, monkeypatch, text: str): + """Point get_config_path at a throwaway artemis.jsonc.""" + config_path = tmp_path / "artemis.jsonc" + config_path.write_text(text, encoding="utf-8") + monkeypatch.setattr( + "artemis.config.paths.get_config_path", + lambda *_args, **_kwargs: config_path, + ) + return config_path + + +@patch.object(ModelService, "_get_llm_provider_and_model", return_value=_CONFIGURED) +def test_get_llm_options_reads_presets_and_the_configured_default( + _configured, tmp_path, monkeypatch +): + _jsonc_config(tmp_path, monkeypatch, _PRESETS_JSONC) + + options = ModelService.get_llm_options() + + # The allowlist is the LLMProvider literal, in declaration order. + assert options["providers"][:3] == ["openai", "google", "openrouter"] + assert "anthropic" in options["providers"] + # Presets keep their jsonc order and drop the rest of the preset body. + assert options["presets"] == [ + {"name": "gemini-flash", "provider": "google", "model": "gemini-2.5-flash"}, + {"name": "local-ollama", "provider": "openai", "model": "llama3.2-vision"}, + ] + # The default comes from the validated LLM config, not the raw block. + assert options["default"] == {"provider": "google", "model": "gemini-3.8-flash"} + + +@patch.object(ModelService, "_get_llm_provider_and_model", return_value=_CONFIGURED) +def test_get_llm_options_degrades_to_no_presets_when_the_config_is_unreadable( + _configured, tmp_path, monkeypatch +): + def _missing(*_args, **_kwargs): + raise FileNotFoundError("artemis.jsonc") + + monkeypatch.setattr("artemis.config.paths.get_config_path", _missing) + + options = ModelService.get_llm_options() + + assert options["presets"] == [] + assert options["default"] == {"provider": "google", "model": "gemini-3.8-flash"} + + +@patch.object(ModelService, "_get_llm_provider_and_model", return_value=_CONFIGURED) +def test_get_llm_options_ignores_a_presets_block_that_is_not_a_mapping( + _configured, tmp_path, monkeypatch +): + _jsonc_config(tmp_path, monkeypatch, '{"presets": ["nope"]}') + + assert ModelService.get_llm_options()["presets"] == [] + + +def test_get_llm_options_real_config_parses_without_coupling_to_content(): + """Smoke: the shipped artemis.jsonc parses; asserts shape, not content.""" + from artemis.config.constants import LLMProvider + + options = ModelService.get_llm_options() + + assert len(options["presets"]) >= 1 + assert all({"name", "provider", "model"} <= set(preset) for preset in options["presets"]) + assert {preset["provider"] for preset in options["presets"]} <= set(get_args(LLMProvider)) + assert options["default"]["provider"] and options["default"]["model"] + + +@pytest.mark.asyncio +@patch.object(ModelService, "_get_llm_provider_and_model", return_value=_CONFIGURED) +async def test_llm_options_route_returns_the_service_payload(_configured, tmp_path, monkeypatch): + _jsonc_config(tmp_path, monkeypatch, _PRESETS_JSONC) + + payload = await tasks_router.get_llm_options() + + assert set(payload) == {"providers", "presets", "default"} + assert payload["default"] == {"provider": "google", "model": "gemini-3.8-flash"} diff --git a/tests/unit/admin_console/test_session_llm_echo.py b/tests/unit/admin_console/test_session_llm_echo.py new file mode 100644 index 00000000..39f2394a --- /dev/null +++ b/tests/unit/admin_console/test_session_llm_echo.py @@ -0,0 +1,189 @@ +# 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. + +"""Server-side producer/consumer contract for the per-task LLM override echo. + +The SDK writes ``llm_model`` / ``llm_provider`` into the session's schemaless +``device_info``; the console reads them back and reports the model the run +actually used. These tests drive the real route handler against a real database +row, so a producer/consumer mismatch fails here rather than in the UI. +""" + +import asyncio +import json +from unittest.mock import MagicMock, patch + +import pytest + +from apps.admin_console.database.connection import db_session +from apps.admin_console.database.repositories.session_repository import SessionRepository +from apps.admin_console.routers import sessions as sessions_router +from apps.admin_console.services.model_service import ModelService + +# The configured (global) LLM, patched so the tests never read artemis.jsonc. +_CONFIGURED = ("google", "gemini-3.8-flash") + + +def _insert_session(db_path, session_id, device_info): + with db_session(db_path) as conn: + conn.execute( + "INSERT INTO sessions (session_id, initial_goal, start_time, status, device_info)" + " VALUES (?, ?, ?, ?, ?)", + (session_id, "goal", 1.0, "completed", json.dumps(device_info)), + ) + conn.commit() + + +def _real_repo_row(tmp_path, monkeypatch, session_id, device_info): + """Back the router's repository with a real DB row for one session.""" + db_path = tmp_path / "sessions.db" + repo = SessionRepository(db_path) + _insert_session(db_path, session_id, device_info) + monkeypatch.setattr(sessions_router.session_repo, "get_session_by_id", repo.get_session_by_id) + return repo + + +@patch.object(ModelService, "_get_llm_provider_and_model", return_value=_CONFIGURED) +def test_get_session_details_echoes_the_recorded_override(_configured, tmp_path, monkeypatch): + """GET /api/sessions/{id} reports the model this run was pinned to.""" + _real_repo_row( + tmp_path, + monkeypatch, + "s1", + { + "profile": "pro", + "llm_model": "gpt-5.1", + "llm_provider": "openai", + "run_tuning": {"verification_level": "final"}, + }, + ) + + payload = asyncio.run(sessions_router.get_session_details("s1")) + + assert payload["llm_model"] == "gpt-5.1" + assert payload["llm_provider"] == "openai" + # Both halves sit at the top level of the payload: that is the contract + # artemis-client's TaskResult.from_payload reads, so the echo resolves for a + # remote caller and not only against a mocked transport. + assert isinstance(payload["llm_model"], str) + assert isinstance(payload["llm_provider"], str) + # Parity with the list endpoint: model_info reflects the same override. + assert payload["model_info"]["id"] == "gpt-5.1" + assert payload["model_info"]["provider"] == "openai" + assert payload["model_info"]["name"] == "Pro" + # The stored row is untouched. + assert json.loads(payload["device_info"])["llm_model"] == "gpt-5.1" + + +@patch.object(ModelService, "_get_llm_provider_and_model", return_value=_CONFIGURED) +def test_get_session_details_reports_null_for_pre_override_rows(_configured, tmp_path, monkeypatch): + """A row written before the override existed reports null, not a guess.""" + _real_repo_row(tmp_path, monkeypatch, "s2", {"profile": "pro"}) + + payload = asyncio.run(sessions_router.get_session_details("s2")) + + assert payload["llm_model"] is None + assert payload["llm_provider"] is None + assert payload["model_info"]["id"] == "gemini-3.8-flash" + assert payload["model_info"]["provider"] == "google" + + +@pytest.mark.asyncio +@patch.object(ModelService, "_get_llm_provider_and_model", return_value=_CONFIGURED) +async def test_list_sessions_model_info_prefers_the_recorded_override(_configured, monkeypatch): + """The list endpoint applies the same override to every row.""" + repo = MagicMock() + repo.get_all_sessions.return_value = [ + { + "session_id": "session-override", + "status": "completed", + "start_time": 1.0, + "device_info": json.dumps( + {"profile": "pro", "llm_model": "gpt-5.1", "llm_provider": "openai"} + ), + }, + { + "session_id": "session-legacy", + "status": "completed", + "start_time": 2.0, + "device_info": json.dumps({"profile": "pro"}), + }, + ] + repo.get_video_recordings_map.return_value = {} + repo.get_latest_video_recordings_map.return_value = {} + repo.get_agent_trace_names_map.return_value = {} + repo.get_llm_traces_for_profiles_map.return_value = {} + + monkeypatch.setattr(sessions_router, "session_repo", repo, raising=False) + monkeypatch.setattr( + sessions_router.media_service, "build_video_index", MagicMock(return_value={}) + ) + monkeypatch.setattr( + sessions_router.media_service, "resolve_video_url", MagicMock(return_value=None) + ) + + rows = {row["session_id"]: row for row in await sessions_router.list_sessions()} + + assert rows["session-override"]["model_info"]["id"] == "gpt-5.1" + assert rows["session-override"]["model_info"]["provider"] == "openai" + assert rows["session-legacy"]["model_info"]["id"] == "gemini-3.8-flash" + assert rows["session-legacy"]["model_info"]["provider"] == "google" + + +@pytest.mark.asyncio +@patch.object(ModelService, "_get_llm_provider_and_model", return_value=_CONFIGURED) +async def test_list_sessions_reports_the_override_for_a_profile_less_row(_configured, monkeypatch): + """No resolvable profile must not hide a stored override. + + Regression: the row used to fall back to the global ``default_model_info`` + whenever the profile stayed unresolved, dropping the pinned model. + """ + repo = MagicMock() + repo.get_all_sessions.return_value = [ + { + "session_id": "session-profileless-override", + "status": "completed", + "start_time": 1.0, + "device_info": json.dumps({"llm_model": "gpt-5.1", "llm_provider": "openai"}), + }, + { + "session_id": "session-profileless-legacy", + "status": "completed", + "start_time": 2.0, + "device_info": None, + }, + ] + repo.get_video_recordings_map.return_value = {} + repo.get_latest_video_recordings_map.return_value = {} + repo.get_agent_trace_names_map.return_value = {} + repo.get_llm_traces_for_profiles_map.return_value = {} + + monkeypatch.setattr(sessions_router, "session_repo", repo, raising=False) + monkeypatch.setattr( + sessions_router.media_service, "build_video_index", MagicMock(return_value={}) + ) + monkeypatch.setattr( + sessions_router.media_service, "resolve_video_url", MagicMock(return_value=None) + ) + + rows = {row["session_id"]: row for row in await sessions_router.list_sessions()} + + override = rows["session-profileless-override"]["model_info"] + assert (override["id"], override["provider"]) == ("gpt-5.1", "openai") + # The architecture stays Flash until the trace-name pass settles the profile. + assert override["name"] == "Flash" + # No keys at all: the untouched global default, built once per request. + legacy = rows["session-profileless-legacy"]["model_info"] + assert (legacy["id"], legacy["provider"]) == ("gemini-3.8-flash", "google") + assert legacy == sessions_router.model_service.get_active_model_info() diff --git a/tests/unit/admin_console/test_task_device_readiness.py b/tests/unit/admin_console/test_task_device_readiness.py index 3ecc9ef4..9ba44919 100644 --- a/tests/unit/admin_console/test_task_device_readiness.py +++ b/tests/unit/admin_console/test_task_device_readiness.py @@ -17,6 +17,7 @@ from fastapi import HTTPException import pytest +from pydantic import ValidationError from apps.admin_console.routers import tasks from apps.admin_console.schemas.task_schema import RunRequest @@ -226,3 +227,82 @@ async def test_explicit_device_proceeds_when_enumeration_is_indeterminate(monkey enqueue_tasks.assert_awaited_once() _, kwargs = enqueue_tasks.call_args assert kwargs.get("device_serial") == "pixel-10" + + +@pytest.mark.asyncio +async def test_run_task_forwards_per_task_llm_override(monkeypatch): + """The per-task LLM override reaches enqueue_tasks untouched.""" + unlocked_probe = ProbeResult( + id="android_adb", + category=ProbeCategory.DEVICE, + title="Device / Emulator Connected", + status=ProbeStatus.PASS, + is_blocker=True, + summary="Connected", + description="Ready.", + ) + run_probe = AsyncMock(return_value=unlocked_probe) + enqueue_tasks = AsyncMock(return_value={"status": "started", "tasks": []}) + monkeypatch.setattr(tasks.readiness_engine, "run_device_submission_probe", run_probe) + monkeypatch.setattr(tasks.task_queue_service, "enqueue_tasks", enqueue_tasks) + + await tasks.run_task( + RunRequest(goal="Audit checkout", llm_model=" gpt-5.1 ", llm_provider=" OpenAI ") + ) + + _, kwargs = enqueue_tasks.call_args + # The schema trims the model and normalises the provider to lower case. + assert kwargs.get("llm_model") == "gpt-5.1" + assert kwargs.get("llm_provider") == "openai" + + +@pytest.mark.asyncio +async def test_run_task_without_override_leaves_llm_override_unset(monkeypatch): + unlocked_probe = ProbeResult( + id="android_adb", + category=ProbeCategory.DEVICE, + title="Device / Emulator Connected", + status=ProbeStatus.PASS, + is_blocker=True, + summary="Connected", + description="Ready.", + ) + run_probe = AsyncMock(return_value=unlocked_probe) + enqueue_tasks = AsyncMock(return_value={"status": "started", "tasks": []}) + monkeypatch.setattr(tasks.readiness_engine, "run_device_submission_probe", run_probe) + monkeypatch.setattr(tasks.task_queue_service, "enqueue_tasks", enqueue_tasks) + + await tasks.run_task(RunRequest(goal="Open Settings", llm_model=" ")) + + _, kwargs = enqueue_tasks.call_args + assert kwargs.get("llm_model") is None + assert kwargs.get("llm_provider") is None + + +def test_run_request_rejects_unknown_llm_provider(): + """An unusable provider is rejected at the API boundary, not mid-task.""" + with pytest.raises(ValidationError): + RunRequest(goal="Open Settings", llm_provider="acme-cloud") + + +def test_run_request_rejects_provider_without_model(): + """A provider alone would pin nothing, so FastAPI answers 422.""" + with pytest.raises(ValidationError, match="llm_provider requires llm_model"): + RunRequest(goal="Open Settings", llm_provider="openai") + + +def test_run_request_accepts_blank_provider_with_model(): + """Blank keeps the blank -> None semantics, so no 422 and no override.""" + request = RunRequest(goal="Open Settings", llm_model="gpt-5.1", llm_provider=" ") + + assert request.llm_model == "gpt-5.1" + assert request.llm_provider is None + + +def test_run_request_accepts_flash_profile_with_override(): + """The Flash execution profile and the LLM override are independent knobs.""" + request = RunRequest(goal="Open Settings", profile="flash", llm_model="gemini-3.8-pro") + + assert request.profile == "flash" + assert request.llm_model == "gemini-3.8-pro" + assert request.llm_provider is None diff --git a/tests/unit/admin_console/test_task_queue_service.py b/tests/unit/admin_console/test_task_queue_service.py index 24a572cd..77f39a75 100644 --- a/tests/unit/admin_console/test_task_queue_service.py +++ b/tests/unit/admin_console/test_task_queue_service.py @@ -23,6 +23,7 @@ from apps.admin_console.core.state import state from apps.admin_console.routers.tasks import get_status +from apps.admin_console.services.model_service import ModelService from apps.admin_console.services.task_queue_service import TaskQueueService, task_queue_service from artemis.runtime.device_lock import DeviceLockOwner from artemis.runtime.adb_endpoint import AdbEndpoint @@ -292,6 +293,8 @@ async def fake_subprocess_exec(*args, **kwargs): app_path="/path/to/app.apk", verification_level=" Checkpoints ", explorer_mode="ULTRA", + llm_model=" gemini-3.8-pro ", + llm_provider="OpenAI", ) for _ in range(30): @@ -314,8 +317,14 @@ async def fake_subprocess_exec(*args, **kwargs): # to the worker as the CLI's existing spelling. assert cmd[cmd.index("--verification-level") + 1] == "checkpoints" assert cmd[cmd.index("--explorer-pro-mode") + 1] == "ultra" + # The per-task LLM override reaches the worker as --model/--provider; + # the model keeps its spelling, the provider is lower-cased. + assert cmd[cmd.index("--model") + 1] == "gemini-3.8-pro" + assert cmd[cmd.index("--provider") + 1] == "openai" assert enqueue_result["tasks"][0]["verification_level"] == "checkpoints" assert enqueue_result["tasks"][0]["explorer_mode"] == "ultra" + assert enqueue_result["tasks"][0]["llm_model"] == "gemini-3.8-pro" + assert enqueue_result["tasks"][0]["llm_provider"] == "openai" assert ( executed_kwargs[0]["env"]["ARTEMIS_DEVICE_QUEUE_TICKET"] == (enqueue_result["tasks"][0]["queue_ticket"]) @@ -557,6 +566,8 @@ async def test_status_reports_external_global_owner_without_ipc_connection(): ): repo.get_latest_session.return_value = None repo.get_session_by_id.return_value = None + # No session row was needed, so the model echo stays unset. + models.resolve_session_llm_override.return_value = (None, None) models.get_active_model_info.return_value = None result = await get_status() @@ -566,6 +577,241 @@ async def test_status_reports_external_global_owner_without_ipc_connection(): assert result["goal"] == "CLI task: inspect settings" +@pytest.mark.asyncio +@patch.object( + ModelService, "_get_llm_provider_and_model", return_value=("google", "gemini-3.8-flash") +) +async def test_get_status_reports_the_stored_override_with_a_known_profile(_configured): + """A pinned model still wins when the profile came from the worker state. + + Regression: the row used to be loaded only when ``active_profile`` was + falsy, so a worker that set ``current_profile`` hid the override and the + status showed the global model. + """ + state.current_profile = "pro" + state.active_session_id = "sess-1" + with ( + patch.object(TaskQueueService, "ensure_worker_running"), + patch("apps.admin_console.routers.tasks.session_repo") as repo, + ): + repo.get_latest_session.return_value = {"session_id": "sess-1"} + repo.get_session_by_id.return_value = { + "session_id": "sess-1", + "status": "running", + "device_info": json.dumps( + {"profile": "pro", "llm_model": "gpt-5.1", "llm_provider": "openai"} + ), + } + repo.get_background_tasks.return_value = [] + + result = await get_status() + + assert result["model_info"]["id"] == "gpt-5.1" + assert result["model_info"]["provider"] == "openai" + assert result["model_info"]["name"] == "Pro" + # One row read, reused for the model echo, and the expensive trace lookups + # stay skipped while a profile is already known. + repo.get_session_by_id.assert_called_once_with("sess-1") + repo.get_llm_traces_for_profile.assert_not_called() + repo.get_agent_trace_names.assert_not_called() + + +@pytest.mark.asyncio +@patch.object( + ModelService, "_get_llm_provider_and_model", return_value=("google", "gemini-3.8-flash") +) +async def test_get_status_keeps_the_global_model_for_legacy_rows(_configured): + """A row without the override keys reports the configured model.""" + state.current_profile = "pro" + state.active_session_id = "sess-1" + with ( + patch.object(TaskQueueService, "ensure_worker_running"), + patch("apps.admin_console.routers.tasks.session_repo") as repo, + ): + repo.get_latest_session.return_value = {"session_id": "sess-1"} + repo.get_session_by_id.return_value = { + "session_id": "sess-1", + "status": "running", + "device_info": json.dumps({"profile": "pro"}), + } + repo.get_background_tasks.return_value = [] + + result = await get_status() + + assert result["model_info"]["id"] == "gemini-3.8-flash" + assert result["model_info"]["provider"] == "google" + assert result["model_info"]["name"] == "Pro" + + +@pytest.mark.asyncio +@patch.object( + ModelService, "_get_llm_provider_and_model", return_value=("google", "gemini-3.8-flash") +) +async def test_get_status_queue_still_carries_the_llm_override(_configured): + """Queued tickets keep reporting the model they will run on.""" + state.current_profile = "pro" + state.queue_items = [ + { + "session_id": "queued-1", + "goal": "Goal A", + "status": "pending", + "llm_model": "gpt-5.1", + "llm_provider": "openai", + } + ] + with ( + patch.object(TaskQueueService, "ensure_worker_running"), + patch("apps.admin_console.routers.tasks.session_repo") as repo, + ): + repo.get_latest_session.return_value = None + repo.get_background_tasks.return_value = [] + + result = await get_status() + + (queued,) = result["queue"] + assert queued["llm_model"] == "gpt-5.1" + assert queued["llm_provider"] == "openai" + # No session row to read here, so the status falls back to the global model. + repo.get_session_by_id.assert_not_called() + assert result["model_info"]["id"] == "gemini-3.8-flash" + + +@pytest.mark.asyncio +@patch.object( + ModelService, "_get_llm_provider_and_model", return_value=("google", "gemini-3.8-flash") +) +async def test_get_status_connection_branch_reports_the_stored_override(_configured): + """A tracked connection reports the pinned model, not the global one. + + Regression: this second return path built ``conn_model_info`` from the + profile alone, so a task tracked through ``active_connections`` showed the + global model even when the run had an override. + """ + state.current_profile = "pro" + state.active_connections["sess-1"] = { + "goal": "Goal A", + "profile": "pro", + "pid": None, + } + with ( + patch.object(TaskQueueService, "ensure_worker_running"), + patch("apps.admin_console.routers.tasks.session_repo") as repo, + ): + repo.get_latest_session.return_value = {"session_id": "sess-1"} + repo.get_session_by_id.return_value = { + "session_id": "sess-1", + "status": "running", + "device_info": json.dumps( + {"profile": "pro", "llm_model": "gpt-5.1", "llm_provider": "openai"} + ), + } + repo.get_background_tasks.return_value = [] + + result = await get_status() + + assert result["session_id"] == "sess-1" + assert result["model_info"]["id"] == "gpt-5.1" + assert result["model_info"]["provider"] == "openai" + assert result["model_info"]["name"] == "Pro" + # Reused the row the main path already read, and stayed off the trace + # lookups entirely. + repo.get_session_by_id.assert_called_once_with("sess-1") + repo.get_llm_traces_for_profile.assert_not_called() + repo.get_agent_trace_names.assert_not_called() + + +@pytest.mark.asyncio +@patch.object( + ModelService, "_get_llm_provider_and_model", return_value=("google", "gemini-3.8-flash") +) +async def test_get_status_connection_branch_keeps_the_global_model_for_legacy_rows(_configured): + """A connection without stored override keys keeps reporting the config.""" + state.current_profile = "pro" + state.active_connections["sess-1"] = {"goal": "Goal A", "profile": "pro", "pid": None} + with ( + patch.object(TaskQueueService, "ensure_worker_running"), + patch("apps.admin_console.routers.tasks.session_repo") as repo, + ): + repo.get_latest_session.return_value = {"session_id": "sess-1"} + repo.get_session_by_id.return_value = { + "session_id": "sess-1", + "status": "running", + "device_info": json.dumps({"profile": "pro"}), + } + repo.get_background_tasks.return_value = [] + + result = await get_status() + + assert result["model_info"]["id"] == "gemini-3.8-flash" + assert result["model_info"]["provider"] == "google" + assert result["model_info"]["name"] == "Pro" + + +@pytest.mark.asyncio +@patch.object( + ModelService, "_get_llm_provider_and_model", return_value=("google", "gemini-3.8-flash") +) +async def test_get_status_connection_branch_reads_the_connections_own_session(_configured): + """When another session owns the earlier row, this path reads its own. + + ``running_sid`` can point at a different session than the connection being + reported, so the cached row must not be reused blindly. + """ + other_owner = DeviceLockOwner( + pid=24681, + process_created_at=1234.5, + token="other-owner-token", + device_id="emulator-5554", + description="CLI task: other", + acquired_at="2026-08-24T00:00:00+00:00", + session_id="other-session", + ingress="cli", + ) + state.current_profile = "pro" + state.active_connections["sess-1"] = {"goal": "Goal A", "profile": "pro", "pid": None} + with ( + patch.object(TaskQueueService, "ensure_worker_running"), + # A lock owner for another device but no "active" owner: running_sid then + # resolves to that other session while this path reports the connection. + patch( + "apps.admin_console.routers.tasks.DeviceExecutionLock.get_active_owner", + return_value=None, + ), + patch( + "apps.admin_console.routers.tasks.DeviceExecutionLock.get_active_owners", + return_value={other_owner.device_id: other_owner}, + ), + patch("apps.admin_console.routers.tasks.session_repo") as repo, + ): + repo.get_latest_session.return_value = {"session_id": "sess-1"} + repo.get_background_tasks.return_value = [] + repo.get_session_by_id.side_effect = lambda sid: { + "other-session": { + "session_id": "other-session", + "device_info": json.dumps({"profile": "pro", "llm_model": "other-model"}), + }, + "sess-1": { + "session_id": "sess-1", + "device_info": json.dumps( + {"profile": "pro", "llm_model": "gpt-5.1", "llm_provider": "openai"} + ), + }, + }[str(sid)] + + result = await get_status() + + assert result["session_id"] == "sess-1" + assert result["model_info"]["id"] == "gpt-5.1" + assert result["model_info"]["provider"] == "openai" + # The other session's row is not reused for this connection. + assert [call.args[0] for call in repo.get_session_by_id.call_args_list] == [ + "other-session", + "sess-1", + ] + repo.get_llm_traces_for_profile.assert_not_called() + repo.get_agent_trace_names.assert_not_called() + + @pytest.mark.asyncio async def test_cancel_task_triggers_next_pending_task(): executed_goals = [] @@ -765,6 +1011,63 @@ async def test_enqueue_tasks_unified_ingress(): assert task["goal"] == "Test unified goal" +@pytest.mark.asyncio +async def test_enqueue_tasks_normalises_llm_override_and_omits_when_blank(): + """A blank LLM override means 'not requested' and is persisted as None.""" + with ( + patch.object(TaskQueueService, "ensure_worker_running"), + patch( + "apps.admin_console.services.task_queue_service.DeviceExecutionLock.reserve", + return_value="mock-ticket-llm", + ), + ): + res = await task_queue_service.enqueue_tasks( + ["Goal with model override"], + profile="pro", + llm_model=" gpt-5.1 ", + llm_provider=" OpenAI ", + ) + blank = await task_queue_service.enqueue_tasks( + ["Goal without override"], + profile="pro", + llm_model=" ", + llm_provider="", + ) + + task, blank_task = res["tasks"][0], blank["tasks"][0] + assert task["llm_model"] == "gpt-5.1" + assert task["llm_provider"] == "openai" + assert blank_task["llm_model"] is None + assert blank_task["llm_provider"] is None + + +@pytest.mark.asyncio +async def test_enqueue_tasks_rejects_an_unusable_llm_override(): + """The queue runs the shared validator, so a bad override never gets a slot. + + The API schema already refuses these, but the queue is also reachable from + the daemon and MCP, so it must not persist a task that cannot run. + """ + with ( + patch.object(TaskQueueService, "ensure_worker_running") as worker, + patch( + "apps.admin_console.services.task_queue_service.DeviceExecutionLock.reserve", + return_value="mock-ticket-llm", + ), + ): + with pytest.raises(ValueError, match="unknown llm_provider"): + await task_queue_service.enqueue_tasks( + ["Goal"], profile="pro", llm_model="gpt-5.1", llm_provider="acme-cloud" + ) + with pytest.raises(ValueError, match="llm_provider requires llm_model"): + await task_queue_service.enqueue_tasks( + ["Goal"], profile="pro", llm_model=" ", llm_provider="openai" + ) + + assert worker.call_count == 0 + assert state.queue_items == [] + + @pytest.mark.asyncio async def test_queue_worker_notifies_conversation(): """Verify queue_worker calls notify() when conversation_id is attached to task.""" diff --git a/tests/unit/mcp/test_background_task_runner.py b/tests/unit/mcp/test_background_task_runner.py index c0fbf7c8..b963b627 100644 --- a/tests/unit/mcp/test_background_task_runner.py +++ b/tests/unit/mcp/test_background_task_runner.py @@ -39,16 +39,27 @@ async def test_background_agent_initialization_forwards_health_settings(): @pytest.mark.asyncio @pytest.mark.parametrize( - ("knobs", "expect_level", "expect_mode"), + ("knobs", "expect_level", "expect_mode", "expect_override"), [ - ({"verification_level": "strict", "explorer_pro_mode": "ultra"}, "strict", "ultra"), - ({}, None, None), + ( + {"verification_level": "strict", "explorer_pro_mode": "ultra"}, + "strict", + "ultra", + None, + ), + ( + {"llm_model": " gpt-5.1 ", "llm_provider": "OpenAI"}, + None, + None, + {"model": "gpt-5.1", "provider": "openai"}, + ), + ({}, None, None, None), ], ) async def test_run_task_applies_pro_tuning_to_agent_config( - tmp_path, monkeypatch, knobs, expect_level, expect_mode + tmp_path, monkeypatch, knobs, expect_level, expect_mode, expect_override ): - """The detached runner applies --verification-level / --explorer-pro-mode on the builder.""" + """The detached runner applies the Pro knobs and the per-task LLM override.""" from types import SimpleNamespace from unittest.mock import patch @@ -100,5 +111,13 @@ async def test_run_task_applies_pro_tuning_to_agent_config( else: fake_builder.with_verification_level.assert_called_once_with(expect_level) fake_builder.with_explorer.assert_called_once_with(pro_mode=expect_mode) + # The LLM override rides on the task request, not on the shared builder. + task_builder = ( + fake_agent.new_task.return_value.with_name.return_value.with_trace_recording.return_value + ) + if expect_override is None: + task_builder.with_llm_override.assert_not_called() + else: + task_builder.with_llm_override.assert_called_once_with(**expect_override) fake_agent.run_task.assert_awaited_once() assert trace_store.read_status(trace_id)["status"] == "completed" diff --git a/tests/unit/mcp/test_mcp_tools.py b/tests/unit/mcp/test_mcp_tools.py index 0064ed9a..080661b2 100644 --- a/tests/unit/mcp/test_mcp_tools.py +++ b/tests/unit/mcp/test_mcp_tools.py @@ -55,11 +55,20 @@ def test_tool_signatures(): assert "device_serial" in sig_run.parameters assert "verification_level" in sig_run.parameters assert "explorer_mode" in sig_run.parameters + assert "llm_model" in sig_run.parameters + assert "llm_provider" in sig_run.parameters assert sig_run.parameters["verification_level"].default is None assert sig_run.parameters["explorer_mode"].default is None + # `model` already means the Flash/Pro profile, so the LLM override uses + # distinct names. + assert sig_run.parameters["model"].default == "Flash" + assert sig_run.parameters["llm_model"].default is None + assert sig_run.parameters["llm_provider"].default is None # The tool description is the only schema an MCP caller sees. assert "verification_level" in (mobile_run_task.__doc__ or "") assert "explorer_mode" in (mobile_run_task.__doc__ or "") + assert "llm_model" in (mobile_run_task.__doc__ or "") + assert "llm_provider" in (mobile_run_task.__doc__ or "") # mobile_manage_task signature check sig_manage = inspect.signature(mobile_manage_task) @@ -273,6 +282,103 @@ def test_mobile_run_task_forwards_pro_tuning_to_daemon(temp_trace_env, monkeypat assert submit.call_args.kwargs["explorer_mode"] == "pro" +def test_mobile_run_task_forwards_llm_override_to_background_runner(temp_trace_env): + process = MagicMock(pid=779) + with ( + patch("mcp_server.tools.task_runner.DeviceExecutionLock.reserve", return_value="t"), + patch("mcp_server.tools.task_runner.DeviceExecutionLock.transfer_reservation"), + patch("mcp_server.tools.task_runner.subprocess.Popen", return_value=process) as popen, + ): + mobile_run_task( + task_desc="Audit the checkout flow", + model="Pro", + llm_model=" gemini-3.8-pro ", + llm_provider="OpenAI", + ) + + cmd = popen.call_args.args[0] + # Distinct from the runner's own --model flag, which carries the Flash/Pro + # profile; the LLM override uses its own spelling. + assert cmd[cmd.index("--llm-model") + 1] == "gemini-3.8-pro" + assert cmd[cmd.index("--llm-provider") + 1] == "openai" + # The profile flag is untouched by the override. + assert cmd[cmd.index("--model") + 1] == "Pro" + + +def test_mobile_run_task_omits_llm_override_flags_when_unset(temp_trace_env): + process = MagicMock(pid=780) + with ( + patch("mcp_server.tools.task_runner.DeviceExecutionLock.reserve", return_value="t"), + patch("mcp_server.tools.task_runner.DeviceExecutionLock.transfer_reservation"), + patch("mcp_server.tools.task_runner.subprocess.Popen", return_value=process) as popen, + ): + mobile_run_task(task_desc="Open Settings", model="Pro", llm_model=" ") + + cmd = popen.call_args.args[0] + assert "--llm-model" not in cmd + assert "--llm-provider" not in cmd + + +def test_mobile_run_task_rejects_unknown_llm_provider_before_creating_a_trace(temp_trace_env): + with pytest.raises(ValueError, match="unknown llm_provider"): + mobile_run_task(task_desc="Open Settings", model="Pro", llm_provider="acme-cloud") + # Rejected before init_trace: nothing was written to the trace store. + assert os.listdir(temp_trace_env) == [] + + +def test_mobile_run_task_rejects_llm_provider_without_model(temp_trace_env): + with pytest.raises(ValueError, match="requires llm_model"): + mobile_run_task(task_desc="Open Settings", model="Pro", llm_provider="openai") + assert os.listdir(temp_trace_env) == [] + + +def test_mobile_run_task_combines_flash_profile_with_llm_override(temp_trace_env): + """The default Flash profile and the LLM override travel together.""" + process = MagicMock(pid=781) + with ( + patch("mcp_server.tools.task_runner.DeviceExecutionLock.reserve", return_value="t"), + patch("mcp_server.tools.task_runner.DeviceExecutionLock.transfer_reservation"), + patch("mcp_server.tools.task_runner.subprocess.Popen", return_value=process) as popen, + ): + mobile_run_task( + task_desc="Open Settings", + model="Flash", + llm_model="gemini-3.8-pro", + llm_provider="openai", + ) + + cmd = popen.call_args.args[0] + assert cmd[cmd.index("--model") + 1] == "Flash" + assert cmd[cmd.index("--llm-model") + 1] == "gemini-3.8-pro" + assert cmd[cmd.index("--llm-provider") + 1] == "openai" + + +def test_mobile_run_task_forwards_llm_override_to_daemon(temp_trace_env, monkeypatch): + monkeypatch.delenv("ARTEMIS_STANDALONE", raising=False) + with ( + patch( + "mcp_server.tools.task_runner.ensure_daemon_running", + return_value=(True, "http://127.0.0.1:8000"), + ), + patch( + "mcp_server.tools.task_runner.submit_task_to_daemon", + return_value={"status": "started", "tasks": [{"session_id": "daemon-sid-3"}]}, + ) as submit, + patch("mcp_server.tools.task_runner.subprocess.Popen") as popen, + ): + result = mobile_run_task( + task_desc="Audit via Daemon", + model="Pro", + llm_model="gpt-5.1", + llm_provider="openai", + ) + + popen.assert_not_called() + assert result["trace_id"] == "daemon-sid-3" + assert submit.call_args.kwargs["llm_model"] == "gpt-5.1" + assert submit.call_args.kwargs["llm_provider"] == "openai" + + def test_mobile_manage_task_unknown_trace(temp_trace_env): res = mobile_manage_task(action="status", trace_id="non-existent-trace") assert res["status"] == "unknown" diff --git a/tests/unit/runtime/test_daemon_client.py b/tests/unit/runtime/test_daemon_client.py index 67a90d3b..3a326224 100644 --- a/tests/unit/runtime/test_daemon_client.py +++ b/tests/unit/runtime/test_daemon_client.py @@ -183,6 +183,49 @@ def test_submit_batch_to_daemon_forwards_pro_tuning_knobs(): assert data["explorer_mode"] == "pro" +def test_submit_task_to_daemon_forwards_per_task_llm_override(): + mock_resp = _daemon_ok_response(b'{"status": "queued", "tasks": [{"session_id": "s1"}]}') + with patch("urllib.request.urlopen", return_value=mock_resp) as mock_urlopen: + submit_task_to_daemon( + "audit goal", + profile="pro", + llm_model="gpt-5.1", + llm_provider="openai", + ) + data = json.loads(mock_urlopen.call_args[0][0].data.decode("utf-8")) + # JSON field names match RunRequest on /api/run. + assert data["llm_model"] == "gpt-5.1" + assert data["llm_provider"] == "openai" + + +def test_submit_task_to_daemon_llm_override_defaults_to_null(): + mock_resp = _daemon_ok_response(b'{"status": "queued", "tasks": [{"session_id": "s1"}]}') + with patch("urllib.request.urlopen", return_value=mock_resp) as mock_urlopen: + submit_task_to_daemon("plain goal", profile="flash") + data = json.loads(mock_urlopen.call_args[0][0].data.decode("utf-8")) + # The daemon client keeps sending nulls verbatim (RunRequest defaults to None). + assert data["llm_model"] is None + assert data["llm_provider"] is None + + +def test_submit_batch_to_daemon_forwards_per_task_llm_override(): + from artemis.runtime.daemon_client import submit_batch_to_daemon + + mock_resp = _daemon_ok_response( + b'{"status": "queued", "tasks": [{"session_id": "a"}, {"session_id": "b"}]}' + ) + with patch("urllib.request.urlopen", return_value=mock_resp) as mock_urlopen: + submit_batch_to_daemon( + ["goal a", "goal b"], + profile="pro", + llm_model="gemini-3.8-pro", + llm_provider="google", + ) + data = json.loads(mock_urlopen.call_args[0][0].data.decode("utf-8")) + assert data["llm_model"] == "gemini-3.8-pro" + assert data["llm_provider"] == "google" + + def test_stop_task_on_daemon(): from artemis.runtime.daemon_client import stop_task_on_daemon diff --git a/tests/unit/services/test_task_llm_override.py b/tests/unit/services/test_task_llm_override.py new file mode 100644 index 00000000..fdd0cee0 --- /dev/null +++ b/tests/unit/services/test_task_llm_override.py @@ -0,0 +1,249 @@ +# 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. + +"""Per-task LLM override: request plumbing and endpoint resolution precedence. + +Precedence is task override > ``artemis.jsonc`` node config > built-in default, +and the override lives on the per-task context, so nothing global is mutated. +""" + +from types import SimpleNamespace + +import pytest + +import artemis.sdk.agent as agent_module +from artemis.config.llm import LLM, LLMConfig, LLMConfigUtils, LLMWithFallback +from artemis.llm.router import ModelProvider +from artemis.sdk.agent import Agent +from artemis.sdk.builders.task_request_builder import TaskRequestBuilder +from artemis.sdk.types.agent import AgentConfig, ServerConfig +from artemis.sdk.types.task import AgentProfile, TaskRequestCommon +from artemis.services.llm import _resolve_endpoint + + +def _node(provider: str = "google", model: str = "gemini-3.8-flash", **extra): + return SimpleNamespace(provider=provider, model=model, temperature=0.0, **extra) + + +def _ctx(node, llm_model=None, llm_provider=None): + """Minimal stand-in for the per-task ArtemisContext.""" + return SimpleNamespace( + llm_config=SimpleNamespace( + get_agent=lambda name: node, + get_utils=lambda name: node, + ), + llm_model=llm_model, + llm_provider=llm_provider, + ) + + +def _real_llm_config() -> LLMConfig: + """A real LLMConfig (not a stub) whose operator node is gemini-3.8-flash.""" + node = LLMWithFallback( + provider="google", + model="gemini-3.8-flash", + temperature=0.0, + fallback=LLM(provider="google", model="gemini-3.7-flash", temperature=0.0), + ) + return LLMConfig( + planner=node, + utils=LLMConfigUtils(outputter=node, hopper=node), + summarizer=node, + operator=node, + operator_summarizer=node, + log_reader_sub_agent=node, + log_analyzer=node, + diagnoser=node, + checker=node, + planner_avatar=node, + history_analyzer_expert=node, + diagnoser_expert=node, + explorer=node, + ) + + +def test_override_wins_over_node_config(): + endpoint = _resolve_endpoint( + _ctx(_node(), llm_model="gpt-5.1", llm_provider="openai"), + "operator", + ) + + assert endpoint.provider == ModelProvider.OPENAI + assert endpoint.model_name == "gpt-5.1" + + +def test_node_config_used_when_no_override(): + endpoint = _resolve_endpoint(_ctx(_node(provider="anthropic", model="claude-x")), "operator") + + assert endpoint.provider == ModelProvider.ANTHROPIC + assert endpoint.model_name == "claude-x" + + +def test_model_only_override_keeps_node_provider(): + endpoint = _resolve_endpoint( + _ctx(_node(provider="google"), llm_model="gemini-3.8-pro"), "planner" + ) + + assert endpoint.provider == ModelProvider.GOOGLE + assert endpoint.model_name == "gemini-3.8-pro" + + +def test_blank_override_falls_back_to_node_config(): + endpoint = _resolve_endpoint(_ctx(_node(), llm_model=" ", llm_provider=""), "planner") + + assert endpoint.provider == ModelProvider.GOOGLE + assert endpoint.model_name == "gemini-3.8-flash" + + +def test_override_also_applies_to_resolved_fallback(): + node = _node(fallback=_node(model="gemini-3.7-flash")) + + endpoint = _resolve_endpoint( + _ctx(node, llm_model="gpt-5.1", llm_provider="openai"), + "operator", + use_fallback=True, + ) + + assert endpoint.provider == ModelProvider.OPENAI + assert endpoint.model_name == "gpt-5.1" + + +def test_override_does_not_mutate_the_shared_llm_config(): + """The override is per task: the shared LLMConfig is only read.""" + config = _real_llm_config() + + def ctx(**override): + return SimpleNamespace(llm_config=config, **override) + + endpoint = _resolve_endpoint( + ctx(llm_model="gpt-5.1", llm_provider="openai"), + "operator", + ) + + assert (endpoint.provider, endpoint.model_name) == (ModelProvider.OPENAI, "gpt-5.1") + # The node this task overrode is untouched, fallback included. + assert (config.operator.provider, config.operator.model) == ("google", "gemini-3.8-flash") + assert (config.operator.fallback.provider, config.operator.fallback.model) == ( + "google", + "gemini-3.7-flash", + ) + # Another node and a later task without an override still see the original. + assert _resolve_endpoint(ctx(), "checker").model_name == "gemini-3.8-flash" + # A second task may pin a different model on another node at the same time. + other = _resolve_endpoint(ctx(llm_model="claude-x", llm_provider="anthropic"), "explorer") + assert (other.provider, other.model_name) == (ModelProvider.ANTHROPIC, "claude-x") + assert config.operator is config.checker + + +def test_builder_rejects_provider_without_model(): + """A provider alone pins nothing, so it is refused at build time.""" + with pytest.raises(ValueError, match="llm_provider requires llm_model"): + TaskRequestBuilder(goal="Audit checkout").with_llm_override(provider="openai") + + +def test_builder_rejects_blank_model_with_provider(): + builder = TaskRequestBuilder(goal="Audit checkout") + + with pytest.raises(ValueError, match="llm_provider requires llm_model"): + builder.with_llm_override(model=" ", provider="openai") + + +def test_builder_rejects_unknown_provider(): + with pytest.raises(ValueError, match="unknown llm_provider"): + TaskRequestBuilder(goal="Audit checkout").with_llm_override( + model="gpt-5.1", provider="NotAProvider" + ) + + +def test_builder_carries_override_on_the_task_request(): + builder = TaskRequestBuilder(goal="Audit checkout") + + request = builder.with_llm_override(model=" gpt-5.1 ", provider=" OpenAI ").build() + + assert request.llm_model == "gpt-5.1" + assert request.llm_provider == "openai" + + +def test_builder_blank_override_leaves_request_unset(): + builder = TaskRequestBuilder(goal="Open Settings") + + request = builder.with_llm_override(model=" ").build() + + assert request.llm_model is None + assert request.llm_provider is None + + +def _agent_config() -> AgentConfig: + profile = AgentProfile(name="default", llm_config=_real_llm_config()) + return AgentConfig( + agent_profiles={"default": profile}, + task_request_defaults=TaskRequestCommon(goal="Audit checkout"), + default_profile=profile, + servers=ServerConfig(adb_host="127.0.0.1", adb_port=5037), + ) + + +def _device_data_for(tmp_path, monkeypatch, request) -> dict: + """Run the real ``_prepare_tracing`` and return the device_info it stored. + + Only the DataEngine sink is stubbed: it would otherwise create a session in + the local database. The device_info assembly itself is the real code. + """ + recorded: dict = {} + + class RecordingDataEngine: + def __init__(self, ctx): + recorded["ctx"] = ctx + + def start_session(self, goal, device_info=None, session_id=None): + recorded["goal"] = goal + recorded["device_info"] = device_info + + monkeypatch.setattr(agent_module, "DataEngine", RecordingDataEngine) + agent = SimpleNamespace(_tmp_traces_dir=tmp_path, _config=_agent_config(), _session_id=None) + task = SimpleNamespace(request=request, get_name=lambda: "task-1") + context = SimpleNamespace(device=None, execution_setup=None, data_engine=None) + + Agent._prepare_tracing(agent, task, context) + + return recorded["device_info"] + + +def test_device_data_records_the_llm_override(tmp_path, monkeypatch): + """The server producer: the override lands in the schemaless device_info.""" + request = ( + TaskRequestBuilder(goal="Audit checkout") + .using_profile("pro") + .with_llm_override(model=" gpt-5.1 ", provider=" OpenAI ") + .build() + ) + + device_data = _device_data_for(tmp_path, monkeypatch, request) + + assert device_data["llm_model"] == "gpt-5.1" + assert device_data["llm_provider"] == "openai" + # Alongside the pre-existing echoes, which the console reads the same way. + assert device_data["profile"] == "pro" + assert device_data["run_tuning"] + + +def test_device_data_omits_the_llm_keys_without_an_override(tmp_path, monkeypatch): + """Legacy behaviour: no override means no keys, so old and new rows agree.""" + request = TaskRequestBuilder(goal="Open Settings").using_profile("flash").build() + + device_data = _device_data_for(tmp_path, monkeypatch, request) + + assert "llm_model" not in device_data + assert "llm_provider" not in device_data + assert device_data["profile"] == "flash" diff --git a/tests/unit/test_cli.py b/tests/unit/test_cli.py index 11e9077a..1dcf1ffb 100644 --- a/tests/unit/test_cli.py +++ b/tests/unit/test_cli.py @@ -48,6 +48,64 @@ def test_cli_run_help(): assert "--traces-path" in result.output assert "--verification-level" in result.output assert "--explorer-pro-mode" in result.output + assert "--model" in result.output + assert "--provider" in result.output + + +def test_cli_run_forwards_per_task_llm_override_in_standalone_mode(monkeypatch): + """`artemis run --standalone --model/--provider` threads the override into execute_task.""" + import artemis.interfaces.cli.commands.run as run_module + + captured: dict = {} + + async def fake_execute_task(**kwargs): + captured.update(kwargs) + + monkeypatch.setattr(run_module, "execute_task", fake_execute_task) + # Cosmetic device-status display would otherwise hit a real ADB server. + monkeypatch.setattr(run_module, "display_device_status", lambda *a, **k: None) + monkeypatch.setenv("ARTEMIS_STANDALONE", "1") + result = runner.invoke( + app, + [ + "run", + "--standalone", + "--model", + "gpt-5.1", + "--provider", + "openai", + "Open Settings", + ], + ) + assert result.exit_code == 0, result.output + assert captured["llm_model"] == "gpt-5.1" + assert captured["llm_provider"] == "openai" + + +def test_cli_run_forwards_per_task_llm_override_to_daemon(monkeypatch): + """Daemon-routed runs carry the override as /api/run JSON fields.""" + import artemis.runtime as runtime + + monkeypatch.delenv("ARTEMIS_STANDALONE", raising=False) + captured: dict = {} + + def fake_submit_task(**kwargs): + captured.update(kwargs) + return {"status": "queued", "tasks": [{"session_id": "sid-1"}]} + + monkeypatch.setattr( + runtime, "ensure_daemon_running", lambda **_kw: (True, "http://127.0.0.1:8000") + ) + monkeypatch.setattr(runtime, "submit_task_to_daemon", fake_submit_task) + monkeypatch.setattr(runtime, "wait_for_daemon_task", lambda *_, **__: {"status": "completed"}) + + result = runner.invoke( + app, + ["run", "--model", "gpt-5.1", "--provider", "openai", "Open Settings"], + ) + assert result.exit_code == 0, result.output + assert captured["llm_model"] == "gpt-5.1" + assert captured["llm_provider"] == "openai" def test_cli_batch_help(): @@ -58,6 +116,8 @@ def test_cli_batch_help(): assert "--delay" in result.output assert "--verification-level" in result.output assert "--explorer-pro-mode" in result.output + assert "--model" in result.output + assert "--provider" in result.output def test_cli_batch_forwards_pro_tuning_in_standalone_mode(monkeypatch): @@ -153,6 +213,105 @@ def test_run_batch_tasks_applies_pro_tuning_to_agent_config(monkeypatch): fake_agent.run_task.assert_awaited_once_with(goal="Goal A", profile="pro") +def test_cli_batch_forwards_llm_override_in_standalone_mode(monkeypatch): + """`artemis batch --standalone --model/--provider` reaches run_batch_tasks.""" + import artemis.interfaces.cli.commands.batch as batch_module + + captured: dict = {} + + async def fake_run_batch_tasks(tasks, **kwargs): + captured["tasks"] = tasks + captured.update(kwargs) + + monkeypatch.setattr(batch_module, "run_batch_tasks", fake_run_batch_tasks) + result = runner.invoke( + app, + [ + "batch", + "--standalone", + "--profile", + "flash", + "--model", + "gemini-3.8-pro", + "--provider", + "openai", + "Open Settings", + ], + ) + assert result.exit_code == 0, result.output + assert captured["tasks"] == ["Open Settings"] + assert captured["llm_model"] == "gemini-3.8-pro" + assert captured["llm_provider"] == "openai" + + +def test_cli_batch_forwards_llm_override_to_daemon(monkeypatch): + """Daemon-routed batches carry the override as /api/run JSON fields.""" + import artemis.runtime as runtime + + monkeypatch.delenv("ARTEMIS_STANDALONE", raising=False) + captured: dict = {} + + def fake_submit_batch(goals, **kwargs): + captured["goals"] = goals + captured.update(kwargs) + return {"tasks": [{"session_id": "sid-1", "goal": goals[0]}]} + + monkeypatch.setattr(runtime, "ensure_daemon_running", lambda **_: (True, "http://x:1")) + monkeypatch.setattr(runtime, "submit_batch_to_daemon", fake_submit_batch) + monkeypatch.setattr(runtime, "wait_for_daemon_task", lambda *_, **__: {"status": "completed"}) + + result = runner.invoke( + app, + ["batch", "--model", "gemini-3.8-pro", "--provider", "openai", "Goal A"], + ) + assert result.exit_code == 0, result.output + assert captured["goals"] == ["Goal A"] + assert captured["llm_model"] == "gemini-3.8-pro" + assert captured["llm_provider"] == "openai" + + +def test_run_batch_tasks_pins_llm_override_on_every_goal(monkeypatch): + """The batch override reaches each goal's own request and nothing global.""" + import asyncio + from unittest.mock import AsyncMock, MagicMock + + import artemis.interfaces.cli.commands.batch as batch_module + from artemis.sdk.builders.task_request_builder import TaskRequestBuilder + + fake_builder = MagicMock() + fake_builders = MagicMock() + fake_builders.AgentConfig.with_default_profile.return_value = fake_builder + fake_agent = MagicMock() + fake_agent.init = AsyncMock() + fake_agent.run_task = AsyncMock(return_value="ok") + fake_agent.clean = AsyncMock() + fake_agent.new_task = MagicMock(side_effect=lambda goal: TaskRequestBuilder(goal=goal)) + + monkeypatch.setattr(batch_module, "initialize_llm_config", lambda: MagicMock()) + monkeypatch.setattr(batch_module, "AgentProfile", MagicMock()) + monkeypatch.setattr(batch_module, "Builders", fake_builders) + monkeypatch.setattr(batch_module, "Agent", MagicMock(return_value=fake_agent)) + + asyncio.run( + batch_module.run_batch_tasks( + ["Goal A", "Goal B"], + profile_name="flash", + delay_seconds=0, + llm_model="gemini-3.8-pro", + llm_provider="openai", + ) + ) + + assert fake_agent.run_task.await_count == 2 + requests = [call.kwargs["request"] for call in fake_agent.run_task.await_args_list] + assert [request.goal for request in requests] == ["Goal A", "Goal B"] + assert all(request.profile == "flash" for request in requests) + assert all( + (request.llm_model, request.llm_provider) == ("gemini-3.8-pro", "openai") + for request in requests + ) + + def test_cli_trace_help(): """Verify 'artemis trace --help' lists trace subcommands.""" result = runner.invoke(app, ["trace", "--help"]) diff --git a/tests/unit/test_llm_override_normalization.py b/tests/unit/test_llm_override_normalization.py new file mode 100644 index 00000000..5d07b097 --- /dev/null +++ b/tests/unit/test_llm_override_normalization.py @@ -0,0 +1,130 @@ +# 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. + +"""The single normaliser every per-task LLM override goes through. + +``normalize_llm_override`` replaced four copies of the same strip/validate +shape (API schema, SDK builder, MCP tool + background runner, queue service). +These tests pin the semantics once, then pin that the entry points really share +it, so a fourth copy cannot creep back in with its own wording. +""" + +import pytest +from pydantic import ValidationError + +from apps.admin_console.schemas.task_schema import RunRequest +from artemis.config.llm_override import SUPPORTED_LLM_PROVIDERS, normalize_llm_override +from artemis.sdk.builders.task_request_builder import TaskRequestBuilder + + +def test_model_is_stripped_but_never_re_casefolded(): + """Model ids are provider-defined and case-sensitive, so only trim them.""" + assert normalize_llm_override(" gpt-5.1 ") == ("gpt-5.1", None) + assert normalize_llm_override("GPT-5.1") == ("GPT-5.1", None) + + +def test_provider_is_stripped_and_lower_cased(): + assert normalize_llm_override("gpt-5.1", " OpenAI ") == ("gpt-5.1", "openai") + # Every supported provider survives its own canonical spelling. + for provider in SUPPORTED_LLM_PROVIDERS: + assert normalize_llm_override("m", provider.upper()) == ("m", provider) + + +def test_blank_values_mean_not_requested(): + assert normalize_llm_override(None, None) == (None, None) + assert normalize_llm_override("", "") == (None, None) + assert normalize_llm_override(" ", " ") == (None, None) + # A model on its own is a valid override (the provider then stays put). + assert normalize_llm_override(" gpt-5.1 ", " ") == ("gpt-5.1", None) + + +def test_provider_without_a_model_is_rejected(): + with pytest.raises(ValueError, match="llm_provider requires llm_model"): + normalize_llm_override(None, "openai") + # A blank model is the same mistake: it pins nothing. + with pytest.raises(ValueError, match="llm_provider requires llm_model"): + normalize_llm_override(" ", "openai") + + +def test_unknown_provider_is_rejected_and_lists_the_supported_ones(): + with pytest.raises(ValueError) as excinfo: + normalize_llm_override("gpt-5.1", "acme-cloud") + + message = str(excinfo.value) + assert "unknown llm_provider 'acme-cloud'" in message + for provider in SUPPORTED_LLM_PROVIDERS: + assert provider in message + + +def test_error_prefix_lands_in_both_messages(): + with pytest.raises(ValueError, match="^builder: unknown llm_provider"): + normalize_llm_override("gpt-5.1", "nope", error_prefix="builder: ") + with pytest.raises(ValueError, match="^builder: llm_provider requires llm_model"): + normalize_llm_override(None, "openai", error_prefix="builder: ") + + +def test_non_string_input_is_stringified(): + """JSON callers can hand us a number; it is used as its own identifier. + + The MCP tool used to coerce and the schema used to reject, so coercion is + the unifying choice: an unusable value now degrades to a string id instead + of failing (or crashing) somewhere further downstream. + """ + assert normalize_llm_override(5, "openai") == ("5", "openai") # type: ignore[arg-type] + # A non-string provider is stringified and then still validated. + with pytest.raises(ValueError, match="unknown llm_provider False"): + normalize_llm_override("gpt-5.1", False) # type: ignore[arg-type] + + +# --- The entry points share the one implementation --------------------------- + + +def test_api_schema_normalises_through_the_shared_helper(): + request = RunRequest(goal="Audit checkout", llm_model=" gpt-5.1 ", llm_provider=" OpenAI ") + + assert (request.llm_model, request.llm_provider) == ("gpt-5.1", "openai") + assert RunRequest(goal="Audit checkout", llm_model=" ").llm_model is None + + +@pytest.mark.parametrize( + ("kwargs", "expected"), + [ + ({"llm_provider": "acme-cloud", "llm_model": "gpt-5.1"}, "unknown llm_provider"), + ({"llm_provider": "openai"}, "llm_provider requires llm_model"), + ({"llm_model": " ", "llm_provider": "openai"}, "llm_provider requires llm_model"), + ], +) +def test_api_schema_rejects_the_same_inputs_as_the_helper(kwargs, expected): + with pytest.raises(ValidationError, match=expected): + RunRequest(goal="Audit checkout", **kwargs) + + +def test_builder_reports_the_shared_message_for_the_same_inputs(): + """The SDK says the same thing, prefixed with the method that failed.""" + with pytest.raises(ValueError, match="^with_llm_override: unknown llm_provider"): + TaskRequestBuilder(goal="Audit checkout").with_llm_override( + model="gpt-5.1", provider="acme-cloud" + ) + with pytest.raises(ValueError, match="^with_llm_override: llm_provider requires llm_model"): + TaskRequestBuilder(goal="Audit checkout").with_llm_override(provider="openai") + + +def test_builder_stores_what_the_helper_returns(): + request = ( + TaskRequestBuilder(goal="Audit checkout") + .with_llm_override(model=" gpt-5.1 ", provider=" OpenAI ") + .build() + ) + + assert (request.llm_model, request.llm_provider) == ("gpt-5.1", "openai")