diff --git a/README.md b/README.md index 63db7f505..753ade008 100644 --- a/README.md +++ b/README.md @@ -99,6 +99,7 @@ Some examples require extra dependencies. See each sample's directory for specif without wrapping them in a workflow. * [open_telemetry](open_telemetry) - Trace workflows with OpenTelemetry. * [openai_agents](openai_agents) - Run OpenAI Agents SDK agents as durable Temporal workflows. +* [openrouter](openrouter) - Call OpenRouter from Activities: fan out a prompt batch, and pause instead of failing when the budget or credits run out. * [patching](patching) - Alter workflows safely with `patch` and `deprecate_patch`. * [polling](polling) - Recommended implementation of an activity that needs to periodically poll an external resource waiting its successful completion. * [prometheus](prometheus) - Configure Prometheus metrics on clients/workers. diff --git a/openai_agents/model_providers/README.md b/openai_agents/model_providers/README.md index df8f286e8..c9b939864 100644 --- a/openai_agents/model_providers/README.md +++ b/openai_agents/model_providers/README.md @@ -30,6 +30,23 @@ The example uses Anthropic Claude by default but can be modified to use other Li Find more LiteLLM providers at: https://docs.litellm.ai/docs/providers +#### OpenRouter +Uses [OpenRouter](https://openrouter.ai/) as the model provider, so the agent can run on any of the hundreds of models OpenRouter serves through one API key. OpenRouter speaks the OpenAI Chat Completions API, so the stock `OpenAIProvider` works once it is pointed at OpenRouter's base URL. + +Start the OpenRouter worker: +```bash +export OPENROUTER_API_KEY="your_openrouter_api_key" + +uv run openai_agents/model_providers/run_openrouter_worker.py +``` + +Then run the example in a separate terminal: +```bash +uv run openai_agents/model_providers/run_openrouter_workflow.py +``` + +The workflow uses `openrouter/auto`, so OpenRouter picks a model per request; change `OPENROUTER_MODEL` in [workflows/openrouter_workflow.py](workflows/openrouter_workflow.py) to pin any OpenRouter model slug. Tools that run inside the Workflow must be `async`: the Agents SDK runs sync tools in a thread, which the Workflow sandbox does not allow. See the [openrouter](../../openrouter) sample for calling OpenRouter directly from Activities with cost tracking and budgets. + ### Extra #### GPT-OSS with Ollama diff --git a/openai_agents/model_providers/run_openrouter_worker.py b/openai_agents/model_providers/run_openrouter_worker.py new file mode 100644 index 000000000..498f08109 --- /dev/null +++ b/openai_agents/model_providers/run_openrouter_worker.py @@ -0,0 +1,78 @@ +import asyncio +import logging +import os +from datetime import timedelta + +from agents import OpenAIProvider, set_tracing_disabled +from openai import AsyncOpenAI +from temporalio.client import Client +from temporalio.openai_agents import ModelActivityParameters, OpenAIAgentsPlugin +from temporalio.worker import Worker + +from openai_agents.model_providers.workflows.openrouter_workflow import ( + OpenRouterAgentWorkflow, +) + + +# @@@SNIPSTART python-openai-agents-openrouter-provider +def openrouter_provider() -> OpenAIProvider: + """OpenAI Agents SDK model provider backed by OpenRouter. + + OpenRouter speaks the OpenAI Chat Completions API, so the stock provider + works once it is pointed at OpenRouter's base URL. Client retries are off: + the plugin runs each model call as a Temporal Activity, and Temporal owns + the retries. + """ + default_headers: dict[str, str] = {} + # Optional app attribution for OpenRouter's rankings. + if referer := os.getenv("OPENROUTER_HTTP_REFERER"): + default_headers["HTTP-Referer"] = referer + if title := os.getenv("OPENROUTER_APP_TITLE"): + default_headers["X-OpenRouter-Title"] = title + + client = AsyncOpenAI( + base_url="https://openrouter.ai/api/v1", + api_key=os.environ["OPENROUTER_API_KEY"], + max_retries=0, + default_headers=default_headers or None, + ) + # These samples use Chat Completions, OpenRouter's primary endpoint; the + # Agents SDK defaults to the Responses API, which OpenRouter also offers. + return OpenAIProvider(openai_client=client, use_responses=False) + + +# @@@SNIPEND + + +async def main(): + # Disable Agents SDK tracing: the default exporter sends traces to OpenAI's + # backend, which needs an OpenAI API key that this sample does not have. + set_tracing_disabled(disabled=True) + + logging.basicConfig(level=logging.WARNING) + logging.getLogger("temporalio.workflow").setLevel(logging.DEBUG) + + client = await Client.connect( + "localhost:7233", + plugins=[ + OpenAIAgentsPlugin( + model_params=ModelActivityParameters( + start_to_close_timeout=timedelta(seconds=60) + ), + model_provider=openrouter_provider(), + ), + ], + ) + + worker = Worker( + client, + task_queue="openai-agents-model-providers-task-queue", + workflows=[ + OpenRouterAgentWorkflow, + ], + ) + await worker.run() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/openai_agents/model_providers/run_openrouter_workflow.py b/openai_agents/model_providers/run_openrouter_workflow.py new file mode 100644 index 000000000..ffa309b32 --- /dev/null +++ b/openai_agents/model_providers/run_openrouter_workflow.py @@ -0,0 +1,29 @@ +import asyncio + +from temporalio.client import Client +from temporalio.openai_agents import OpenAIAgentsPlugin + +from openai_agents.model_providers.workflows.openrouter_workflow import ( + OpenRouterAgentWorkflow, +) + + +async def main(): + client = await Client.connect( + "localhost:7233", + plugins=[ + OpenAIAgentsPlugin(), + ], + ) + + result = await client.execute_workflow( + OpenRouterAgentWorkflow.run, + "What's the weather in Tokyo?", + id="openai-agents-openrouter-workflow-id", + task_queue="openai-agents-model-providers-task-queue", + ) + print(f"Result: {result}") + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/openai_agents/model_providers/workflows/openrouter_workflow.py b/openai_agents/model_providers/workflows/openrouter_workflow.py new file mode 100644 index 000000000..696350e0e --- /dev/null +++ b/openai_agents/model_providers/workflows/openrouter_workflow.py @@ -0,0 +1,30 @@ +from __future__ import annotations + +from agents import Agent, Runner, function_tool +from temporalio import workflow + +# OpenRouter's Auto Router picks a model per request and takes tool support +# into account. Any OpenRouter model slug works here instead. +OPENROUTER_MODEL = "openrouter/auto" + + +@workflow.defn +class OpenRouterAgentWorkflow: + @workflow.run + async def run(self, prompt: str) -> str: + # Tools that run inside the Workflow must be async: the Agents SDK runs + # sync tools in a thread, which the Workflow sandbox does not allow. + @function_tool + async def get_weather(city: str) -> str: + workflow.logger.debug(f"Getting weather for {city}") + return f"The weather in {city} is sunny." + + agent = Agent( + name="Assistant", + instructions="You only respond in haikus. When asked about the weather always use the tool to get the current weather.", + model=OPENROUTER_MODEL, + tools=[get_weather], + ) + + result = await Runner.run(agent, prompt) + return result.final_output diff --git a/openrouter/README.md b/openrouter/README.md new file mode 100644 index 000000000..7d5dc3752 --- /dev/null +++ b/openrouter/README.md @@ -0,0 +1,75 @@ +# OpenRouter + +These samples call [OpenRouter](https://openrouter.ai/) from Temporal Activities. OpenRouter serves hundreds of models from many providers behind one OpenAI-compatible API and one API key, and picks providers and models per request. Temporal handles everything around those calls: retries with backoff, fan-out with bounded concurrency, crash recovery, pausing for a human, and a durable record of each prompt's result, cost, and retry history. + +| Sample | Description | +|--------|-------------| +| [prompt_batch](prompt_batch) | Fan one OpenRouter call out per prompt with OpenRouter's Auto Router, and collect answer, model, and cost per prompt. Shows Temporal-owned retries, `Retry-After` handling, and retries served for free from OpenRouter's response cache. Start here. | +| [budget_gate](budget_gate) | The same batch, but it pauses instead of failing when money runs out, whether a soft budget in the Workflow or OpenRouter refusing the call for lack of credits, and resumes on a `raise_budget` Update. | + +For OpenRouter as the model provider behind the [OpenAI Agents SDK plugin](../openai_agents), see [openai_agents/model_providers](../openai_agents/model_providers#openrouter). + +## Prerequisites + +1. Follow the [repository prerequisites](../README.md), then install this sample's dependencies: + + ```bash + uv sync --group openrouter + ``` + +2. Start a local dev server with the [Temporal CLI](https://docs.temporal.io/cli): + + ```bash + temporal server start-dev + ``` + +3. Set an [OpenRouter API key](https://openrouter.ai/settings/keys) in the Worker's environment. A few cents of credit is enough for these samples. + + ```bash + export OPENROUTER_API_KEY="sk-or-v1-..." + ``` + + Optional: set `OPENROUTER_HTTP_REFERER` and `OPENROUTER_APP_TITLE` for [app attribution](https://openrouter.ai/docs/app-attribution) in OpenRouter's rankings. + +The API key stays in the Worker process. Prompts, answers, models, and costs go through the Workflow and are recorded in Event History; the key never does. + +## Running a sample + +Each sample has a Worker and a starter. Run them in separate terminals: + +```bash +# Terminal 1 +uv run --group openrouter openrouter/prompt_batch/run_worker.py + +# Terminal 2 +uv run --group openrouter openrouter/prompt_batch/run_workflow.py "Explain retries in one sentence." "Write a haiku about databases." +``` + +## How the Activity calls OpenRouter + +[activities.py](activities.py) uses the `openai` SDK pointed at `https://openrouter.ai/api/v1`, which is the setup OpenRouter documents for OpenAI-compatible clients. OpenRouter-specific fields go in `extra_body`. Four things matter for durable execution: + +- **Temporal owns retries.** The client is created with `max_retries=0`, so every attempt is one HTTP call driven by the Activity retry policy. Event History records the attempt count and the last failure; each attempt's model, cost, and cache status is logged by the Worker. If you use OpenRouter's official `openrouter` package instead, pass `retry_config=RetryConfig("none", ...)`: by default it retries 5xx and connection errors for up to an hour, invisibly. +- **Errors are classified.** 408, 429, and 5xx raise a retryable `ApplicationError`; 400, 401, 403 (moderation or permissions), and other 4xx raise a non-retryable one. Running out of money gets its own type, `OpenRouterOutOfCredits`: a 402 (OpenRouter documents this for both the account and the API key, with `error.metadata.limit_source` saying which) or, as we have seen a per-key limit return in practice, a 403 `Key limit exceeded`. The one exception is a 402 from OpenRouter's in-flight budget cap, which is transient and retried after `Retry-After`. A `Retry-After` header becomes the next retry delay. OpenRouter can also return HTTP 200 with an `error` body and no `choices`, or with a partial answer and an `error` on the choice; the Activity checks for both. +- **Retries are free when the first call succeeded.** The Activity sends `X-OpenRouter-Cache: true`, so if a Worker dies after OpenRouter answered but before Temporal recorded the result, the retried, byte-identical request is served from OpenRouter's response cache and billed at $0. Nothing per-attempt goes in the request body, so attempts stay identical. +- **Heartbeats.** The Activity heartbeats so a dead Worker is detected after `heartbeat_timeout` (10s) rather than after the full `start_to_close_timeout`. + +Each result carries the concrete model OpenRouter chose, OpenRouter's reported `usage.cost` (or `None` if a response had none), the generation id, and the cache status. The batch's `reported_cost_usd` sums those final-attempt figures; it is not a bill, since an attempt that was billed but whose response never reached Temporal is not in it. `unknown_cost_count` is how many of those results had no cost. + +## What Temporal does and does not guarantee + +Activities are at-least-once. If a Worker dies mid-call, the retry re-sends the request; within the cache TTL that retry costs nothing, but two identical requests in flight at the same time both miss the cache and both bill. Completed Activities are never re-run, so a restarted batch resumes at the first unfinished prompt. + +OpenRouter decides which provider and model serve a request, in milliseconds (Auto Router, `models` fallback lists, provider preferences). Temporal decides what happens over time: waiting out a rate limit, surviving a Worker crash, pausing for hours until a human acts, and keeping the audit trail. + +## Batch size + +Each Activity adds a few events to the Workflow's Event History, and every answer is part of the Workflow result. These samples cap a batch at 100 prompts. For larger batches, use one Workflow per slice, or the pattern in [batch_sliding_window](../batch_sliding_window) with continue-as-new. + +## Tests + +The tests replace OpenRouter with a fake HTTP transport and the Activity with a fake, so they need no API key and make no network calls: + +```bash +uv run --group openrouter pytest tests/openrouter +``` diff --git a/openrouter/__init__.py b/openrouter/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/openrouter/activities.py b/openrouter/activities.py new file mode 100644 index 000000000..e395efc94 --- /dev/null +++ b/openrouter/activities.py @@ -0,0 +1,258 @@ +import asyncio +import json +import os +from datetime import datetime, timedelta, timezone +from email.utils import parsedate_to_datetime +from typing import Any, Mapping, NoReturn, Optional + +from openai import APIStatusError, AsyncOpenAI +from temporalio import activity +from temporalio.exceptions import ApplicationError + +from openrouter.shared import ( + OPENROUTER_BASE_URL, + OpenRouterRequest, + OpenRouterResult, +) + + +def build_client(api_key: Optional[str] = None) -> AsyncOpenAI: + """OpenAI SDK client pointed at OpenRouter. + + Client-side retries are disabled so that Temporal owns every retry: the + attempt count and last failure land in Event History, and each attempt is + logged below. (OpenRouter's official SDKs retry 5xx + and connection errors for up to an hour by default; if you use one of them + instead, turn that off too.) + """ + default_headers: dict[str, str] = {} + # App attribution is optional. When set, OpenRouter lists your app in its + # public rankings; send X-OpenRouter-App-Visibility: hidden on the first + # request to create the app entry as hidden. + if referer := os.getenv("OPENROUTER_HTTP_REFERER"): + default_headers["HTTP-Referer"] = referer + if title := os.getenv("OPENROUTER_APP_TITLE"): + default_headers["X-OpenRouter-Title"] = title + return AsyncOpenAI( + base_url=OPENROUTER_BASE_URL, + api_key=api_key or os.environ["OPENROUTER_API_KEY"], + max_retries=0, + timeout=60.0, + default_headers=default_headers or None, + ) + + +def error_type(status: int) -> str: + """Error type recorded in Event History for an OpenRouter HTTP status.""" + return f"OpenRouterHTTP{status}" + + +# Raised instead of an HTTP status type when the call failed for lack of money: +# a 402 (OpenRouter documents this for both the account and the API key, with +# error.metadata.limit_source saying which) or, as observed in practice, a 403 +# "Key limit exceeded" for a per-key limit. A Workflow can pause on this and +# resume once someone tops up. +OUT_OF_CREDITS = "OpenRouterOutOfCredits" + +# A 402 from OpenRouter's in-flight budget cap is transient: wait for +# Retry-After and try again. +TRANSIENT_402_LIMIT_SOURCE = "openrouter_in_flight_budget" + +# Longest Retry-After the Activity will pass through as the next retry delay. +MAX_RETRY_AFTER = timedelta(minutes=5) + + +def _retry_after(headers: Mapping[str, str]) -> Optional[timedelta]: + """Parse Retry-After in either its delta-seconds or HTTP-date form.""" + value = (headers.get("retry-after") or "").strip() + if not value: + return None + delay: Optional[timedelta] = None + try: + delay = timedelta(seconds=float(value)) + except (ValueError, OverflowError): + try: + delay = parsedate_to_datetime(value) - datetime.now(timezone.utc) + except (TypeError, ValueError, OverflowError): + return None + if delay <= timedelta(0): + return None + # Honor the server, within reason: next_retry_delay overrides the retry + # policy's interval, so cap it rather than park a prompt for hours. + return min(delay, MAX_RETRY_AFTER) + + +def raise_for_status( + status: int, error: Mapping[str, Any], headers: Mapping[str, str] +) -> NoReturn: + """Turn an OpenRouter error into an ApplicationError with the right retry posture. + + `error` is OpenRouter's error object ({"code", "message", "metadata"}). + Retryable: 408 (timeout), 429 (rate limited, honoring Retry-After), any + 5xx (500, 502 model down, 503 no provider available, 524, 529), and the + transient in-flight-budget 402. Non-retryable: other 4xx. 400 is a bad + request, 401 a bad key, 403 a moderation or permission block. Retrying + those only costs time. Out of money is its own type (OUT_OF_CREDITS). + """ + message = str(error.get("message") or "") + metadata = error.get("metadata") + limit_source = metadata.get("limit_source") if isinstance(metadata, dict) else None + transient_402 = status == 402 and limit_source == TRANSIENT_402_LIMIT_SOURCE + if not transient_402 and ( + status == 402 or (status == 403 and "limit exceeded" in message.lower()) + ): + raise ApplicationError( + f"OpenRouter returned HTTP {status}: {message}", + {"status": status}, + type=OUT_OF_CREDITS, + non_retryable=True, + ) + retryable = transient_402 or status in (408, 429) or status >= 500 + raise ApplicationError( + f"OpenRouter returned HTTP {status}: {message}", + {"status": status}, + type=error_type(status), + non_retryable=not retryable, + next_retry_delay=_retry_after(headers) if retryable else None, + ) + + +def _error_object(body: Any) -> dict[str, Any]: + """OpenRouter's error object from either a raw body or openai's APIError.body. + + A raw response body wraps it as {"error": {...}}; openai's APIError.body is + already the inner object. Accept both. + """ + if not isinstance(body, dict): + return {} + inner = body.get("error", body) + return inner if isinstance(inner, dict) else {} + + +def _error_code(error: Mapping[str, Any], default: int) -> int: + code = error.get("code") + return code if isinstance(code, int) else default + + +def _content_to_text(content: Any) -> str: + if isinstance(content, str): + return content + if isinstance(content, list): + return "\n".join( + part["text"] + for part in content + if isinstance(part, dict) and isinstance(part.get("text"), str) + ) + return "" + + +async def _heartbeat_forever(interval: timedelta) -> None: + while True: + await asyncio.sleep(interval.total_seconds()) + activity.heartbeat(activity.info().attempt) + + +class OpenRouterActivities: + def __init__(self, client: AsyncOpenAI) -> None: + self._client = client + + # @@@SNIPSTART python-openrouter-call-activity + @activity.defn + async def call_openrouter(self, request: OpenRouterRequest) -> OpenRouterResult: + """One chat completion. One HTTP call per attempt; Temporal retries.""" + # Heartbeat so a killed Worker is noticed after heartbeat_timeout + # rather than after the full start_to_close_timeout. + heartbeat_timeout = activity.info().heartbeat_timeout + heartbeat_task = ( + asyncio.create_task(_heartbeat_forever(heartbeat_timeout / 2)) + if heartbeat_timeout + else None + ) + try: + return await self._send(request) + finally: + if heartbeat_task: + heartbeat_task.cancel() + + async def _send(self, request: OpenRouterRequest) -> OpenRouterResult: + extra_body: dict[str, Any] = {} + if request.fallback_models: + # OpenRouter tries these in order within the same request. + extra_body["models"] = request.fallback_models + elif request.model == "openrouter/auto": + extra_body["plugins"] = [ + {"id": "auto-router", "cost_tier": request.cost_tier} + ] + model = request.fallback_models[0] if request.fallback_models else request.model + + try: + raw = await self._client.chat.completions.with_raw_response.create( + model=model, + messages=[{"role": "user", "content": request.prompt}], + extra_body=extra_body or None, + extra_headers={ + # Ask OpenRouter to cache the successful response. A retry + # of the byte-identical request within the TTL is served + # from cache and billed at $0. + "X-OpenRouter-Cache": "true", + "X-OpenRouter-Cache-TTL": str(request.cache_ttl_seconds), + }, + ) + except APIStatusError as e: + error = _error_object(e.body) + error.setdefault("message", e.message) + raise_for_status(e.status_code, error, e.response.headers) + # Connection errors and timeouts propagate as-is: Temporal retries them. + + payload = json.loads(raw.text) + if isinstance(payload.get("error"), dict): + # OpenRouter can return HTTP 200 with an error body and no choices + # when the upstream provider failed after the request was accepted. + error = _error_object(payload) + raise_for_status(_error_code(error, 500), error, raw.headers) + choices = payload.get("choices") or [] + choice_error = choices[0].get("error") if choices else None + if isinstance(choice_error, dict): + # Or a 200 with a partial answer and the provider's error on the + # choice itself; a partial answer is not an answer. + raise_for_status(_error_code(choice_error, 500), choice_error, raw.headers) + if not choices: + # No error and no answer: treat like a server error and retry. + raise_for_status(500, {"message": "Response has no choices"}, raw.headers) + + usage = payload.get("usage") or {} + cost = usage.get("cost") + if not isinstance(cost, (int, float)): + # OpenRouter reports cost on every response; if it is ever missing, + # say so rather than pretending the call was free. + activity.logger.warning("OpenRouter response has no usage.cost") + result = OpenRouterResult( + prompt=request.prompt, + model=str(payload.get("model", model)), + answer=_content_to_text((choices[0].get("message") or {}).get("content")) + if choices + else "", + cost_usd=float(cost) if isinstance(cost, (int, float)) else None, + generation_id=str(payload.get("id", "")), + cache_status=raw.headers.get("x-openrouter-cache-status", ""), + ) + activity.logger.info( + "OpenRouter call completed: attempt=%d model=%s cost_usd=%s cache=%s id=%s", + activity.info().attempt, + result.model, + "unknown" if result.cost_usd is None else f"{result.cost_usd:.6f}", + result.cache_status or "-", + result.generation_id, + ) + + if request.fail_once_after_call and activity.info().attempt == 1: + # Demo hook: the Worker "crashes" after the response arrived. The + # retry re-sends the identical request and gets a cache hit. + raise ApplicationError( + "Simulated failure after the response was received", + type="SimulatedFailure", + ) + + return result + + # @@@SNIPEND diff --git a/openrouter/budget_gate/README.md b/openrouter/budget_gate/README.md new file mode 100644 index 000000000..046211127 --- /dev/null +++ b/openrouter/budget_gate/README.md @@ -0,0 +1,118 @@ +# Budget gate + +A prompt batch that pauses instead of failing when money runs out, and resumes when a human raises the budget. + +## What this sample demonstrates + +- A soft budget enforced by the Workflow from the cost OpenRouter reports on every response. When the next call would exceed it, the batch parks on `workflow.wait_condition` and stays parked for as long as it takes (hours, days) without a Worker doing anything. +- OpenRouter refusing a call for lack of credits handled the same way: a 402 for the account or the API key, or the 403 `Key limit exceeded` we have seen a per-key limit return in practice. The failing prompt parks instead of failing, and is re-run after the operator tops up. +- A `raise_budget` Update to resume, with a validator that rejects lowering the budget, and a `spend_report` Query showing spend, reservations, the ledger, and which prompts are parked and why. +- Completed prompts are never re-run. A restarted Worker, or a resumed batch, continues from the first unfinished prompt. + +## Running the sample + +Set `OPENROUTER_API_KEY` (see the [parent README](../README.md)), then: + +```bash +# Terminal 1 +uv run --group openrouter openrouter/budget_gate/run_worker.py + +# Terminal 2: a budget small enough to pause after a prompt or two +uv run --group openrouter openrouter/budget_gate/run_workflow.py --budget-usd 0.0002 --estimate-usd 0.0001 --max-concurrency 1 \ + "Define durable execution in one sentence." "Why do LLM calls belong in Activities?" \ + "Name two causes of HTTP 429." "What does a heartbeat timeout detect?" +``` + +The starter prints the Workflow ID and waits. In a third terminal, watch it pause: + +```bash +temporal workflow query -w --type spend_report +``` + +```json +{ + "budget_usd": 0.0002, + "spent_usd": 0.000629, + "reserved_usd": 0, + "completed": 2, + "paused": { "Name two causes of HTTP 429.": "soft_budget_exhausted" }, + "ledger": [ + { "prompt": "Define durable execution in one sentence.", "model": "deepseek/deepseek-v4-flash-0731", "cost_usd": 0.00009605, "cost_known": true, "generation_id": "gen-...", "cache_status": "MISS" }, + { "prompt": "Why do LLM calls belong in Activities?", "model": "deepseek/deepseek-v4-flash-0731", "cost_usd": 0.0005325, "cost_known": true, "generation_id": "gen-...", "cache_status": "MISS" } + ] +} +``` + +The second prompt alone cost more than the whole budget, so the gate closed at the third; see the note on the soft budget below. + +Raise the budget to resume: + +```bash +uv run --group openrouter openrouter/budget_gate/raise_budget.py 0.01 +# or: temporal workflow update execute -w --name raise_budget --input '0.01' +``` + +The starter then prints the completed batch: + +``` +[deepseek/deepseek-v4-flash-0731] $0.000096 cache=MISS Define durable execution in one sentence. +[deepseek/deepseek-v4-flash-0731] $0.000532 cache=MISS Why do LLM calls belong in Activities? +[deepseek/deepseek-v4-flash-0731] $0.000180 cache=MISS Name two causes of HTTP 429. +[deepseek/deepseek-v4-flash-0731] $0.000065 cache=MISS What does a heartbeat timeout detect? + +Reported cost: $0.000873 (what OpenRouter reported on each prompt's final attempt) +``` + +`temporal workflow show -w ` shows the pause as a `TimerStarted` (the approval timeout), then `WorkflowExecutionUpdateAccepted` and `WorkflowExecutionUpdateCompleted` when the budget is raised, `TimerCanceled`, and the remaining Activities. + +### Out of credits at OpenRouter + +Set a credit limit on your API key in the [OpenRouter dashboard](https://openrouter.ai/settings/keys) below what the batch needs, and run with a generous soft budget: + +```bash +uv run --group openrouter openrouter/budget_gate/run_workflow.py --budget-usd 1.0 --max-concurrency 2 \ + "Define durable execution in one sentence." "Why do LLM calls belong in Activities?" \ + "Name two causes of HTTP 429." "What does a heartbeat timeout detect?" +``` + +When OpenRouter refuses the call (a 402, or the `403 Key limit exceeded` we have seen for a per-key limit), the prompt parks with reason `insufficient_credits`: + +```json +{ + "budget_usd": 1, + "spent_usd": 0, + "reserved_usd": 0.002, + "completed": 0, + "paused": { + "Define durable execution in one sentence.": "insufficient_credits", + "Why do LLM calls belong in Activities?": "insufficient_credits" + }, + "ledger": [] +} +``` + +Raise the key's limit in the dashboard, then send `raise_budget` with the current budget value to resume; the parked prompts are re-run: + +```bash +uv run --group openrouter openrouter/budget_gate/raise_budget.py 1.0 +``` + +If nobody raises the budget within `--approval-timeout-seconds` of the batch starting (default one hour, at most 30 days; one deadline shared by every parked prompt), the batch completes with the remaining prompts listed as skipped. + +## What the soft budget does and does not guarantee + +`spent_usd` is the sum of what OpenRouter reported on each prompt's final, successful attempt (or the estimate, if a response carried no cost). It is not a bill: an attempt that was billed but whose response never reached Temporal, such as a Worker crash after the response, is not in it. For actual spend, use OpenRouter's dashboard or `GET /api/v1/key`. + +The cost of a call is only known after the response, so the Workflow reserves `--estimate-usd` per in-flight call and checks `spent + reserved + estimate <= budget` before starting one. Reservations count against the budget, so overshoot is bounded by `max_concurrency` times how far the real cost of a call exceeds the estimate: nothing if the estimate is high enough, a lot if it is far too low. In the run above the estimate was a fifth of what the second prompt cost, which is why that prompt alone blew the budget before the gate closed at the third. To bound the cost of a single call, set `provider.max_price` in the request (see OpenRouter's provider routing docs). The hard cap is the credit limit on the OpenRouter API key, which is what produces the out-of-credits error. + +While parked, prompts keep their concurrency slots, so at most `max_concurrency` prompts show up in `paused`; the rest wait for a slot. Either way nothing spends until the budget is raised. + +## Files + +| File | Description | +|------|-------------| +| [workflow.py](workflow.py) | `BudgetGateWorkflow`: reservation ledger, pause on soft budget or out-of-credits, `raise_budget` Update with validator, `spend_report` Query. | +| [run_worker.py](run_worker.py) | Builds the OpenRouter client once and runs the Worker. | +| [run_workflow.py](run_workflow.py) | Starts a batch with a budget and prints the result. | +| [raise_budget.py](raise_budget.py) | Sends the `raise_budget` Update. | +| [../activities.py](../activities.py) | `call_openrouter`, shared with [prompt_batch](../prompt_batch). | diff --git a/openrouter/budget_gate/__init__.py b/openrouter/budget_gate/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/openrouter/budget_gate/raise_budget.py b/openrouter/budget_gate/raise_budget.py new file mode 100644 index 000000000..b658564d9 --- /dev/null +++ b/openrouter/budget_gate/raise_budget.py @@ -0,0 +1,30 @@ +import asyncio +import sys + +from temporalio.client import Client +from temporalio.envconfig import ClientConfig + +from openrouter.budget_gate.workflow import BudgetGateWorkflow +from openrouter.shared import BatchResult + + +async def main() -> None: + if len(sys.argv) != 3: + print("usage: raise_budget.py ") + raise SystemExit(2) + workflow_id, new_budget = sys.argv[1], float(sys.argv[2]) + + config = ClientConfig.load_client_connect_config() + config.setdefault("target_host", "localhost:7233") + client = await Client.connect(**config) + + handle = client.get_workflow_handle(workflow_id, result_type=BatchResult) + report = await handle.execute_update(BudgetGateWorkflow.raise_budget, new_budget) + print( + f"Budget is now ${report.budget_usd:.6f}; spent ${report.spent_usd:.6f} " + f"across {report.completed} prompts; paused: {report.paused or 'none'}" + ) + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/openrouter/budget_gate/run_worker.py b/openrouter/budget_gate/run_worker.py new file mode 100644 index 000000000..40f115c10 --- /dev/null +++ b/openrouter/budget_gate/run_worker.py @@ -0,0 +1,32 @@ +import asyncio +import logging + +from temporalio.client import Client +from temporalio.envconfig import ClientConfig +from temporalio.worker import Worker + +from openrouter.activities import OpenRouterActivities, build_client +from openrouter.budget_gate.workflow import BudgetGateWorkflow +from openrouter.shared import BUDGET_GATE_TASK_QUEUE + + +async def main() -> None: + logging.basicConfig(level=logging.INFO) + + config = ClientConfig.load_client_connect_config() + config.setdefault("target_host", "localhost:7233") + client = await Client.connect(**config) + + activities = OpenRouterActivities(build_client()) + + worker = Worker( + client, + task_queue=BUDGET_GATE_TASK_QUEUE, + workflows=[BudgetGateWorkflow], + activities=[activities.call_openrouter], + ) + await worker.run() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/openrouter/budget_gate/run_workflow.py b/openrouter/budget_gate/run_workflow.py new file mode 100644 index 000000000..eb753b88d --- /dev/null +++ b/openrouter/budget_gate/run_workflow.py @@ -0,0 +1,94 @@ +import argparse +import asyncio +import uuid + +from temporalio.client import Client +from temporalio.envconfig import ClientConfig + +from openrouter.budget_gate.workflow import BudgetGateWorkflow +from openrouter.shared import ( + BUDGET_GATE_TASK_QUEUE, + DEFAULT_MODEL, + BatchInput, + BudgetGateInput, +) + +DEFAULT_PROMPTS = [ + "Explain retries in one sentence.", + "Write a haiku about databases.", + "Name three uses for embeddings.", + "Summarize eventual consistency in two sentences.", + "What is a task queue?", + "Give one reason to use idempotency keys.", +] + + +async def main() -> None: + parser = argparse.ArgumentParser( + description="Run a prompt batch that pauses when the budget runs out." + ) + parser.add_argument("prompts", nargs="*", default=DEFAULT_PROMPTS) + parser.add_argument("--model", default=DEFAULT_MODEL) + parser.add_argument( + "--budget-usd", + type=float, + default=0.001, + help="Soft budget. The default is small enough to pause a few prompts in.", + ) + parser.add_argument("--estimate-usd", type=float, default=0.0005) + parser.add_argument("--max-concurrency", type=int, default=2) + parser.add_argument("--approval-timeout-seconds", type=int, default=3600) + parser.add_argument( + "--fail-once", + action="store_true", + help="Fail each Activity's first attempt after the response arrives, " + "so the retry shows a cache hit billed at $0.", + ) + args = parser.parse_args() + + config = ClientConfig.load_client_connect_config() + config.setdefault("target_host", "localhost:7233") + client = await Client.connect(**config) + + workflow_id = f"openrouter-budget-gate-{uuid.uuid4()}" + handle = await client.start_workflow( + BudgetGateWorkflow.run, + BudgetGateInput( + batch=BatchInput( + prompts=args.prompts, + model=args.model, + max_concurrency=args.max_concurrency, + fail_once_after_call=args.fail_once, + ), + budget_usd=args.budget_usd, + estimated_cost_usd=args.estimate_usd, + approval_timeout_seconds=args.approval_timeout_seconds, + ), + id=workflow_id, + task_queue=BUDGET_GATE_TASK_QUEUE, + ) + print(f"Started {workflow_id}") + print("While it runs:") + print(f" temporal workflow query -w {workflow_id} --type spend_report") + print( + f" uv run --group openrouter openrouter/budget_gate/raise_budget.py {workflow_id} 0.05" + ) + print("Waiting for the batch to finish...\n", flush=True) + + result = await handle.result() + for r in result.results: + cost = "unknown" if r.cost_usd is None else f"${r.cost_usd:.6f}" + print(f"[{r.model}] {cost} cache={r.cache_status or '-'} {r.prompt}") + for s in result.skipped: + print(f"[skipped: {s.reason}] {s.prompt}") + print( + f"\nReported cost: ${result.reported_cost_usd:.6f} " + "(what OpenRouter reported on each prompt's final attempt)" + ) + if result.unknown_cost_count: + print(f" {result.unknown_cost_count} prompt(s) came back without a cost") + print(f"Inspect: temporal workflow show -w {workflow_id}") + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/openrouter/budget_gate/workflow.py b/openrouter/budget_gate/workflow.py new file mode 100644 index 000000000..ea4d7f16d --- /dev/null +++ b/openrouter/budget_gate/workflow.py @@ -0,0 +1,284 @@ +import asyncio +import math +from datetime import timedelta +from typing import Callable, Union + +from temporalio import workflow +from temporalio.exceptions import ActivityError, ApplicationError, CancelledError + +# The shared dataclasses are passed through the sandbox so that objects the +# Activity returns are the same classes the Workflow compares against. +with workflow.unsafe.imports_passed_through(): + from openrouter.activities import OUT_OF_CREDITS, OpenRouterActivities + from openrouter.shared import ( + BUDGET_TOLERANCE_USD, + MAX_APPROVAL_TIMEOUT_SECONDS, + MAX_PROMPTS_PER_BATCH, + OPENROUTER_RETRY_POLICY, + BatchResult, + BudgetGateInput, + LedgerEntry, + OpenRouterRequest, + OpenRouterResult, + SkippedPrompt, + SpendReport, + ) + +INSUFFICIENT_CREDITS = OUT_OF_CREDITS + + +@workflow.defn +class BudgetGateWorkflow: + """A prompt batch that pauses instead of failing when money runs out. + + Two things can pause it: the soft budget in the input (checked against the + cost OpenRouter reports per response) and OpenRouter itself refusing the + call for lack of credits (402 for the account, 403 "Key limit exceeded" + for the API key). Either way the batch parks until + a `raise_budget` Update arrives, then resumes exactly where it stopped. + Completed prompts are never re-run. + """ + + @workflow.init + def __init__(self, gate: BudgetGateInput) -> None: + # Set the budget and deadline here rather than in run(): an Update sent + # with update-with-start is handled before run() starts, and must see + # (and be allowed to raise) the real budget. + self._budget_usd = gate.budget_usd + # Validate here, not in run(): a bad value must fail the Workflow, not + # raise while computing the deadline below. + if not 0 <= gate.approval_timeout_seconds <= MAX_APPROVAL_TIMEOUT_SECONDS: + raise ApplicationError( + "approval_timeout_seconds must be between 0 and " + f"{MAX_APPROVAL_TIMEOUT_SECONDS}", + non_retryable=True, + ) + # One deadline for the whole batch: every parked prompt waits until this + # moment, not for its own full approval timeout. + self._deadline = workflow.now() + timedelta( + seconds=gate.approval_timeout_seconds + ) + self._spent_usd = 0.0 + self._reserved_usd = 0.0 + # Bumped by every raise_budget Update, so a prompt parked on a + # credits error can tell that the operator acted even if the soft + # budget did not change. + self._budget_version = 0 + self._ledger: list[LedgerEntry] = [] + self._paused: dict[str, str] = {} + + @workflow.run + async def run(self, gate: BudgetGateInput) -> BatchResult: + batch = gate.batch + if len(batch.prompts) > MAX_PROMPTS_PER_BATCH: + raise ApplicationError( + f"Batch has {len(batch.prompts)} prompts; the limit is " + f"{MAX_PROMPTS_PER_BATCH}.", + non_retryable=True, + ) + if batch.max_concurrency < 1: + raise ApplicationError( + "max_concurrency must be at least 1", non_retryable=True + ) + if ( + not math.isfinite(gate.estimated_cost_usd) + or gate.estimated_cost_usd <= 0 + or not math.isfinite(gate.budget_usd) + or gate.budget_usd < 0 + ): + raise ApplicationError( + "estimated_cost_usd must be a positive finite number and " + "budget_usd a non-negative finite number", + non_retryable=True, + ) + semaphore = asyncio.Semaphore(batch.max_concurrency) + try: + outcomes = await asyncio.gather( + *(self._answer(prompt, gate, semaphore) for prompt in batch.prompts) + ) + finally: + # Good hygiene for any Workflow with handlers: do not return while + # a handler is still running. (raise_budget is synchronous, so this + # is always already true here.) + await workflow.wait_condition(workflow.all_handlers_finished) + + results = [o for o in outcomes if isinstance(o, OpenRouterResult)] + skipped = [o for o in outcomes if isinstance(o, SkippedPrompt)] + return BatchResult( + results=results, + skipped=skipped, + reported_cost_usd=round(self._spent_usd, 6), + unknown_cost_count=sum(1 for e in self._ledger if not e.cost_known), + ) + + # @@@SNIPSTART python-openrouter-budget-gate-handlers + @workflow.update + def raise_budget(self, new_budget_usd: float) -> SpendReport: + """Raise the soft budget and wake every parked prompt. + + Send the current budget unchanged to resume after topping up credits + in the OpenRouter dashboard. + """ + self._budget_usd = new_budget_usd + self._budget_version += 1 + return self.spend_report() + + @raise_budget.validator + def validate_raise_budget(self, new_budget_usd: float) -> None: + if not math.isfinite(new_budget_usd): + raise ValueError("The budget must be a finite number.") + if new_budget_usd < self._budget_usd: + raise ValueError( + f"New budget ${new_budget_usd} is below the current budget " + f"${self._budget_usd}; the budget can only go up." + ) + + @workflow.query + def spend_report(self) -> SpendReport: + return SpendReport( + budget_usd=self._budget_usd, + spent_usd=round(self._spent_usd, 6), + reserved_usd=round(self._reserved_usd, 6), + completed=len(self._ledger), + paused=dict(self._paused), + ledger=list(self._ledger), + ) + + # @@@SNIPEND + + async def _answer( + self, prompt: str, gate: BudgetGateInput, semaphore: asyncio.Semaphore + ) -> Union[OpenRouterResult, SkippedPrompt]: + async with semaphore: + if not await self._reserve(prompt, gate.estimated_cost_usd): + return SkippedPrompt(prompt=prompt, reason="soft_budget_exhausted") + try: + while True: + # Snapshot the budget version before the call, not after it + # fails: a raise_budget that lands while this call is in + # flight must count as the top-up this prompt is waiting + # for, not as one it missed. + seen_version = self._budget_version + try: + result = await workflow.execute_activity_method( + OpenRouterActivities.call_openrouter, + OpenRouterRequest( + prompt=prompt, + model=gate.batch.model, + fail_once_after_call=gate.batch.fail_once_after_call, + ), + start_to_close_timeout=timedelta(seconds=90), + heartbeat_timeout=timedelta(seconds=10), + retry_policy=OPENROUTER_RETRY_POLICY, + ) + break + except ActivityError as e: + cause = e.cause + if isinstance(cause, CancelledError): + # The Activity was cancelled (the Workflow is being + # cancelled); that is not a per-prompt failure. + raise + if ( + isinstance(cause, ApplicationError) + and cause.type == OUT_OF_CREDITS + ): + # Out of credits at OpenRouter. Park until the + # operator tops up and sends raise_budget. + if await self._wait_for_more_credits(prompt, seen_version): + continue + return SkippedPrompt( + prompt=prompt, reason="insufficient_credits" + ) + reason = ( + cause.type + if isinstance(cause, ApplicationError) and cause.type + else type(cause).__name__ + ) + workflow.logger.warning( + "Skipping prompt %r: %s", prompt, reason + ) + return SkippedPrompt(prompt=prompt, reason=reason) + # Charge what OpenRouter reported; if it reported nothing, + # charge the estimate rather than treating the call as free. + charged = ( + result.cost_usd + if result.cost_usd is not None + else gate.estimated_cost_usd + ) + cost_known = result.cost_usd is not None + self._spent_usd += charged + self._ledger.append( + LedgerEntry( + prompt=prompt, + model=result.model, + cost_usd=charged, + cost_known=cost_known, + generation_id=result.generation_id, + cache_status=result.cache_status, + ) + ) + return result + finally: + self._reserved_usd -= gate.estimated_cost_usd + + # @@@SNIPSTART python-openrouter-budget-gate-pause + async def _reserve(self, prompt: str, estimate: float) -> bool: + """Reserve `estimate` against the budget, parking until it fits.""" + + def fits() -> bool: + # Tolerance absorbs float accumulation; see BUDGET_TOLERANCE_USD. + return ( + self._spent_usd + self._reserved_usd + estimate + <= self._budget_usd + BUDGET_TOLERANCE_USD + ) + + # Loop rather than check once: when the budget is raised, every parked + # prompt is woken before any of them runs, so each must re-check after + # waking in case an earlier one already took the new headroom. (Each + # re-park starts a new timer for the remaining time; fine at this scale.) + while not fits(): + workflow.logger.info( + "Soft budget reached (spent $%.6f of $%.6f); pausing %r", + self._spent_usd, + self._budget_usd, + prompt, + ) + if not await self._park(prompt, "soft_budget_exhausted", fits): + return False + self._reserved_usd += estimate + return True + + async def _park(self, prompt: str, reason: str, until: Callable[[], bool]) -> bool: + """Durable pause until `until()` holds or the batch deadline passes. + + Survives Worker restarts and can wait for hours. Returns False when the + deadline passed first. + """ + if until(): + # Nothing to wait for (a top-up already landed during the call). + return True + remaining = self._deadline - workflow.now() + if remaining <= timedelta(0): + return False + self._paused[prompt] = reason + try: + await workflow.wait_condition(until, timeout=remaining) + return True + except asyncio.TimeoutError: + return False + finally: + self._paused.pop(prompt, None) + + # @@@SNIPEND + + async def _wait_for_more_credits(self, prompt: str, seen_version: int) -> bool: + """Park until a raise_budget newer than `seen_version` has arrived. + + If one already has, this returns at once and the prompt is re-run. + """ + workflow.logger.info("OpenRouter says out of credits; pausing %r", prompt) + return await self._park( + prompt, + "insufficient_credits", + lambda: self._budget_version > seen_version, + ) diff --git a/openrouter/prompt_batch/README.md b/openrouter/prompt_batch/README.md new file mode 100644 index 000000000..a614ebf1b --- /dev/null +++ b/openrouter/prompt_batch/README.md @@ -0,0 +1,74 @@ +# Prompt batch + +Fan one OpenRouter call out per prompt and collect the answers. + +## What this sample demonstrates + +- One Activity per prompt, run concurrently under a semaphore, so a slow or failing prompt never blocks the others. +- OpenRouter's Auto Router (`openrouter/auto`) choosing a model per prompt, with the chosen model and OpenRouter's reported cost returned for each. +- Temporal-owned retries: 429 and 5xx retry with backoff and honor `Retry-After`; 4xx errors fail fast and the prompt is reported as skipped instead of failing the batch. +- Retries served from OpenRouter's response cache at $0 when the first call already succeeded. + +## Running the sample + +Set `OPENROUTER_API_KEY` (see the [parent README](../README.md)), then: + +```bash +# Terminal 1 +uv run --group openrouter openrouter/prompt_batch/run_worker.py + +# Terminal 2 +uv run --group openrouter openrouter/prompt_batch/run_workflow.py "Explain retries in one sentence." "Write a haiku about databases." +``` + +Output: + +``` +Starting openrouter-prompt-batch-6923ab3d-... + +[deepseek/deepseek-v4-flash-0731] $0.000022 cache=MISS + Q: Explain retries in one sentence. + A: Retries are the automatic re-attempts of a failed operation, often after a delay ... + +[deepseek/deepseek-v4-flash-0731] $0.000525 cache=MISS + Q: Write a haiku about databases. + A: Columns and table, ... + +Reported cost: $0.000547 (what OpenRouter reported on each prompt's final attempt) +Inspect: temporal workflow show -w openrouter-prompt-batch-6923ab3d-... +``` + +### See a retry that costs nothing + +`--fail-once` makes each Activity fail its first attempt *after* OpenRouter has answered, which is what a Worker crash at the wrong moment looks like. The retry re-sends the identical request and OpenRouter serves it from cache: + +```bash +uv run --group openrouter openrouter/prompt_batch/run_workflow.py --fail-once "Explain idempotency in one sentence." +``` + +``` +[deepseek/deepseek-v4-flash-0731] $0.000000 cache=HIT + Q: Explain idempotency in one sentence. + A: Idempotency means that an operation can be applied multiple times, but the result is the same ... + +Reported cost: $0.000000 (what OpenRouter reported on each prompt's final attempt) +``` + +The first attempt was billed; the reported cost only covers what came back on the final attempt, which is the $0 cache hit. OpenRouter's dashboard is the source of truth for actual spend. + +`temporal workflow show -w ` shows the Activity completing on attempt 2 with the simulated failure as its last failure; the Worker log has one line per attempt with model, cost, and cache status. The cache is keyed on your API key and the exact request body, so running the same prompt again within the cache TTL (10 minutes by default here) is also a hit. OpenRouter writes the cache shortly after the response completes; a retry that arrives before that write lands is a `MISS` and is billed, which you may see occasionally with the one-second retry interval used here. + +### Other options + +- `--model `: any OpenRouter model instead of the Auto Router. +- `--max-concurrency N`: how many prompts are in flight at once (default 5). + +## Files + +| File | Description | +|------|-------------| +| [workflow.py](workflow.py) | `PromptBatchWorkflow`: fan-out under a semaphore, per-prompt failure handling, retry policy. | +| [run_worker.py](run_worker.py) | Builds the OpenRouter client once and runs the Worker. | +| [run_workflow.py](run_workflow.py) | Starts a batch and prints answer, model, cost, and cache status per prompt. | +| [../activities.py](../activities.py) | `call_openrouter`: one HTTP call per attempt, error classification, cache headers, heartbeats. | +| [../shared.py](../shared.py) | Dataclasses shared by starter, Workflow, and Activity. | diff --git a/openrouter/prompt_batch/__init__.py b/openrouter/prompt_batch/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/openrouter/prompt_batch/run_worker.py b/openrouter/prompt_batch/run_worker.py new file mode 100644 index 000000000..3f527733d --- /dev/null +++ b/openrouter/prompt_batch/run_worker.py @@ -0,0 +1,34 @@ +import asyncio +import logging + +from temporalio.client import Client +from temporalio.envconfig import ClientConfig +from temporalio.worker import Worker + +from openrouter.activities import OpenRouterActivities, build_client +from openrouter.prompt_batch.workflow import PromptBatchWorkflow +from openrouter.shared import PROMPT_BATCH_TASK_QUEUE + + +async def main() -> None: + logging.basicConfig(level=logging.INFO) + + config = ClientConfig.load_client_connect_config() + config.setdefault("target_host", "localhost:7233") + client = await Client.connect(**config) + + # One OpenRouter client for the Worker's lifetime, shared by every + # concurrent Activity. Reads OPENROUTER_API_KEY from the environment. + activities = OpenRouterActivities(build_client()) + + worker = Worker( + client, + task_queue=PROMPT_BATCH_TASK_QUEUE, + workflows=[PromptBatchWorkflow], + activities=[activities.call_openrouter], + ) + await worker.run() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/openrouter/prompt_batch/run_workflow.py b/openrouter/prompt_batch/run_workflow.py new file mode 100644 index 000000000..6552617fd --- /dev/null +++ b/openrouter/prompt_batch/run_workflow.py @@ -0,0 +1,67 @@ +import argparse +import asyncio +import uuid + +from temporalio.client import Client +from temporalio.envconfig import ClientConfig + +from openrouter.prompt_batch.workflow import PromptBatchWorkflow +from openrouter.shared import DEFAULT_MODEL, PROMPT_BATCH_TASK_QUEUE, BatchInput + +DEFAULT_PROMPTS = [ + "Explain retries in one sentence.", + "Write a haiku about databases.", +] + + +async def main() -> None: + parser = argparse.ArgumentParser( + description="Run a prompt batch through OpenRouter." + ) + parser.add_argument("prompts", nargs="*", default=DEFAULT_PROMPTS) + parser.add_argument("--model", default=DEFAULT_MODEL) + parser.add_argument("--max-concurrency", type=int, default=5) + parser.add_argument( + "--fail-once", + action="store_true", + help="Fail each Activity's first attempt after the response arrives, " + "so the retry shows a cache hit billed at $0.", + ) + args = parser.parse_args() + + config = ClientConfig.load_client_connect_config() + config.setdefault("target_host", "localhost:7233") + client = await Client.connect(**config) + + workflow_id = f"openrouter-prompt-batch-{uuid.uuid4()}" + print(f"Starting {workflow_id}", flush=True) + result = await client.execute_workflow( + PromptBatchWorkflow.run, + BatchInput( + prompts=args.prompts, + model=args.model, + max_concurrency=args.max_concurrency, + fail_once_after_call=args.fail_once, + ), + id=workflow_id, + task_queue=PROMPT_BATCH_TASK_QUEUE, + ) + + for r in result.results: + cost = "unknown" if r.cost_usd is None else f"${r.cost_usd:.6f}" + print(f"\n[{r.model}] {cost} cache={r.cache_status or '-'}") + print(f" Q: {r.prompt}") + print(f" A: {r.answer.strip()}") + for s in result.skipped: + print(f"\n[skipped: {s.reason}] {s.prompt}") + print( + f"\nReported cost: ${result.reported_cost_usd:.6f} " + "(what OpenRouter reported on each prompt's final attempt)" + ) + if result.unknown_cost_count: + print(f" {result.unknown_cost_count} prompt(s) came back without a cost") + print(f"Inspect: temporal workflow show -w {workflow_id}") + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/openrouter/prompt_batch/workflow.py b/openrouter/prompt_batch/workflow.py new file mode 100644 index 000000000..81c16d465 --- /dev/null +++ b/openrouter/prompt_batch/workflow.py @@ -0,0 +1,88 @@ +import asyncio +from datetime import timedelta +from typing import Union + +from temporalio import workflow +from temporalio.exceptions import ActivityError, ApplicationError, CancelledError + +# The shared dataclasses are passed through the sandbox so that objects the +# Activity returns are the same classes the Workflow compares against. +with workflow.unsafe.imports_passed_through(): + from openrouter.activities import OpenRouterActivities + from openrouter.shared import ( + MAX_PROMPTS_PER_BATCH, + OPENROUTER_RETRY_POLICY, + BatchInput, + BatchResult, + OpenRouterRequest, + OpenRouterResult, + SkippedPrompt, + ) + + +@workflow.defn +class PromptBatchWorkflow: + """Fan one OpenRouter call out per prompt and collect the answers.""" + + @workflow.run + async def run(self, batch: BatchInput) -> BatchResult: + if len(batch.prompts) > MAX_PROMPTS_PER_BATCH: + raise ApplicationError( + f"Batch has {len(batch.prompts)} prompts; the limit is " + f"{MAX_PROMPTS_PER_BATCH}. Split it, or see the README for the " + "sliding-window pattern.", + non_retryable=True, + ) + + if batch.max_concurrency < 1: + raise ApplicationError( + "max_concurrency must be at least 1", non_retryable=True + ) + + # @@@SNIPSTART python-openrouter-prompt-batch-fan-out + semaphore = asyncio.Semaphore(batch.max_concurrency) + outcomes = await asyncio.gather( + *(self._answer(prompt, batch, semaphore) for prompt in batch.prompts) + ) + # @@@SNIPEND + + results = [o for o in outcomes if isinstance(o, OpenRouterResult)] + skipped = [o for o in outcomes if isinstance(o, SkippedPrompt)] + return BatchResult( + results=results, + skipped=skipped, + reported_cost_usd=round(sum(r.cost_usd or 0.0 for r in results), 6), + unknown_cost_count=sum(1 for r in results if r.cost_usd is None), + ) + + async def _answer( + self, prompt: str, batch: BatchInput, semaphore: asyncio.Semaphore + ) -> Union[OpenRouterResult, SkippedPrompt]: + async with semaphore: + try: + return await workflow.execute_activity_method( + OpenRouterActivities.call_openrouter, + OpenRouterRequest( + prompt=prompt, + model=batch.model, + fail_once_after_call=batch.fail_once_after_call, + ), + start_to_close_timeout=timedelta(seconds=90), + heartbeat_timeout=timedelta(seconds=10), + retry_policy=OPENROUTER_RETRY_POLICY, + ) + except ActivityError as e: + cause = e.cause + if isinstance(cause, CancelledError): + # The Activity was cancelled (the Workflow is being cancelled); + # that is not a per-prompt failure. + raise + # One bad prompt should not fail the batch. Record why and + # carry on; the caller decides what to do with skipped prompts. + reason = ( + cause.type + if isinstance(cause, ApplicationError) and cause.type + else type(cause).__name__ + ) + workflow.logger.warning("Skipping prompt %r: %s", prompt, reason) + return SkippedPrompt(prompt=prompt, reason=reason) diff --git a/openrouter/shared.py b/openrouter/shared.py new file mode 100644 index 000000000..8777ed37e --- /dev/null +++ b/openrouter/shared.py @@ -0,0 +1,145 @@ +from dataclasses import dataclass, field +from datetime import timedelta +from typing import Optional + +from temporalio.common import RetryPolicy + +OPENROUTER_BASE_URL = "https://openrouter.ai/api/v1" + +# OpenRouter's Auto Router picks a concrete model per request. The response's +# `model` field reports which one it chose. +DEFAULT_MODEL = "openrouter/auto" + +PROMPT_BATCH_TASK_QUEUE = "openrouter-prompt-batch" +BUDGET_GATE_TASK_QUEUE = "openrouter-budget-gate" + +# Each Activity adds a few events to the Workflow's Event History and each +# answer is stored in the Workflow result payload. Keep batches small enough to +# stay well under the history and payload limits; see the README for the +# sliding-window pattern for larger batches. +MAX_PROMPTS_PER_BATCH = 100 + +# A parked batch can wait at most this long (30 days) for a raise_budget. +MAX_APPROVAL_TIMEOUT_SECONDS = 30 * 24 * 3600 + +# Budget comparisons are done on floats that accumulate per-call costs, so +# allow a hair of slack: ten calls at $0.001 must fit a $0.01 budget. +BUDGET_TOLERANCE_USD = 1e-9 + +# Temporal owns retries: 1s, 2s, 4s, ... capped at 60s, five attempts. The +# Activity marks 4xx errors non-retryable and passes OpenRouter's Retry-After +# through as the next retry delay, so this policy only governs the rest. +OPENROUTER_RETRY_POLICY = RetryPolicy( + initial_interval=timedelta(seconds=1), + backoff_coefficient=2.0, + maximum_interval=timedelta(seconds=60), + maximum_attempts=5, +) + + +@dataclass +class OpenRouterRequest: + """One chat completion request. Everything here ends up in the request + body, so keep it free of per-attempt values (attempt number, timestamps): + OpenRouter's response cache keys on the exact body, and a retried attempt + should be byte-identical to the first one.""" + + prompt: str + model: str = DEFAULT_MODEL + # When set, sent as OpenRouter's `models` list and tried in order. This + # replaces the Auto Router. + fallback_models: list[str] = field(default_factory=list) + # Auto Router cost tier: low, medium, high, xhigh, or max. Only used with + # `openrouter/auto`. + cost_tier: str = "low" + # How long OpenRouter keeps a successful response cached so that a retry of + # the identical request is served for free. + cache_ttl_seconds: int = 600 + # Demo hook: fail the first attempt *after* the response arrives, so the + # retry shows a cache hit billed at $0 in Event History. + fail_once_after_call: bool = False + + +@dataclass +class OpenRouterResult: + prompt: str + model: str + answer: str + # What OpenRouter reported for this attempt's response. None if the + # response carried no usage.cost (it always should). + cost_usd: Optional[float] + generation_id: str + # "HIT" or "MISS" from OpenRouter's X-OpenRouter-Cache-Status header, or "" + # when the header is absent. + cache_status: str + + +@dataclass +class SkippedPrompt: + prompt: str + reason: str + + +@dataclass +class BatchInput: + prompts: list[str] + model: str = DEFAULT_MODEL + max_concurrency: int = 5 + fail_once_after_call: bool = False + + +@dataclass +class BatchResult: + results: list[OpenRouterResult] + skipped: list[SkippedPrompt] + # Sum of the cost OpenRouter reported on each prompt's final, successful + # attempt (budget_gate charges its estimate for a response with no cost). + # Attempts that were billed but whose result never reached Temporal (a + # Worker crash after the response, say) are not in here; OpenRouter's + # dashboard or /api/v1/key is the source of truth for spend. + reported_cost_usd: float + # How many successful prompts came back without a cost. When this is not + # zero, reported_cost_usd is a subtotal of the known costs (prompt_batch) + # or includes the estimate for those prompts (budget_gate). + unknown_cost_count: int = 0 + + +@dataclass +class BudgetGateInput: + # The batch to run: prompts, model, concurrency. Same shape as prompt_batch. + batch: BatchInput + # Soft budget enforced by the Workflow from OpenRouter's reported cost. + budget_usd: float + # Reserved per in-flight call before its real cost is known. Reservations + # count against the budget, so overshoot is bounded by + # batch.max_concurrency * max(actual cost - estimate, 0): nothing if the + # estimate is high enough, unbounded if it is far too low. + estimated_cost_usd: float = 0.001 + # How long, from the start of the batch, parked prompts wait for a + # `raise_budget` Update before the batch gives up on them. One deadline is + # shared by the whole batch. + approval_timeout_seconds: int = 3600 + + +@dataclass +class LedgerEntry: + prompt: str + model: str + # Reported cost, or the batch's estimate when the response had no cost. + cost_usd: float + cost_known: bool + generation_id: str + cache_status: str + + +@dataclass +class SpendReport: + budget_usd: float + spent_usd: float + reserved_usd: float + completed: int + # Prompts currently parked, keyed by prompt text (duplicate prompts share + # one entry), with why: "soft_budget_exhausted" or "insufficient_credits" + # (OpenRouter refused the call for lack of credits). + paused: dict[str, str] + ledger: list[LedgerEntry] diff --git a/pyproject.toml b/pyproject.toml index c84d72e41..b360e9a6f 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -82,6 +82,7 @@ openai-agents = [ "temporalio-openai-agents >= 1.0.0", "requests>=2.32.0,<3", ] +openrouter = ["openai>=1.4.0,<3"] pydantic-converter = ["pydantic>=2.10.6,<3"] sentry = ["sentry-sdk>=2.13.0"] strands-agents = [ diff --git a/tests/openrouter/__init__.py b/tests/openrouter/__init__.py new file mode 100644 index 000000000..e69de29bb diff --git a/tests/openrouter/activity_test.py b/tests/openrouter/activity_test.py new file mode 100644 index 000000000..1af812eb8 --- /dev/null +++ b/tests/openrouter/activity_test.py @@ -0,0 +1,335 @@ +import dataclasses +import json +from datetime import timedelta +from typing import Any, Callable + +import httpx +import pytest +from openai import AsyncOpenAI +from temporalio.exceptions import ApplicationError +from temporalio.testing import ActivityEnvironment + +from openrouter.activities import OpenRouterActivities +from openrouter.shared import OPENROUTER_BASE_URL, OpenRouterRequest + +Handler = Callable[[httpx.Request], httpx.Response] + + +def make_activities(handler: Handler) -> OpenRouterActivities: + """Activities backed by a fake OpenRouter; no network, no API key.""" + client = AsyncOpenAI( + base_url=OPENROUTER_BASE_URL, + api_key="test-key", + max_retries=0, + http_client=httpx.AsyncClient(transport=httpx.MockTransport(handler)), + ) + return OpenRouterActivities(client) + + +def completion_body( + answer: str = "Retries repeat a failed call.", + model: str = "openai/gpt-4o-mini", + cost: Any = 0.000123, +) -> dict[str, Any]: + return { + "id": "gen-123", + "object": "chat.completion", + "created": 0, + "model": model, + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": {"role": "assistant", "content": answer}, + } + ], + "usage": { + "prompt_tokens": 5, + "completion_tokens": 7, + "total_tokens": 12, + "cost": cost, + }, + } + + +async def test_success_returns_model_cost_and_cache_status() -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response( + 200, + json=completion_body(), + headers={"X-OpenRouter-Cache-Status": "MISS"}, + ) + + result = await ActivityEnvironment().run( + make_activities(handler).call_openrouter, + OpenRouterRequest(prompt="Explain retries in one sentence."), + ) + + assert result.model == "openai/gpt-4o-mini" + assert result.answer == "Retries repeat a failed call." + assert result.cost_usd == pytest.approx(0.000123) + assert result.generation_id == "gen-123" + assert result.cache_status == "MISS" + + # Exactly one HTTP call per attempt: the client does not retry on its own. + assert len(requests) == 1 + body = json.loads(requests[0].content) + assert body["model"] == "openrouter/auto" + assert body["plugins"] == [{"id": "auto-router", "cost_tier": "low"}] + assert requests[0].headers["X-OpenRouter-Cache"] == "true" + assert requests[0].headers["X-OpenRouter-Cache-TTL"] == "600" + + +async def test_fallback_models_replace_auto_router() -> None: + requests: list[httpx.Request] = [] + + def handler(request: httpx.Request) -> httpx.Response: + requests.append(request) + return httpx.Response(200, json=completion_body(model="b/second")) + + result = await ActivityEnvironment().run( + make_activities(handler).call_openrouter, + OpenRouterRequest(prompt="hi", fallback_models=["a/first", "b/second"]), + ) + + body = json.loads(requests[0].content) + assert body["model"] == "a/first" + assert body["models"] == ["a/first", "b/second"] + assert "plugins" not in body + assert result.model == "b/second" + + +async def test_rate_limit_is_retryable_and_honors_retry_after() -> None: + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 429, + json={"error": {"code": 429, "message": "Rate limited"}}, + headers={"Retry-After": "7"}, + ) + + with pytest.raises(ApplicationError) as excinfo: + await ActivityEnvironment().run( + make_activities(handler).call_openrouter, OpenRouterRequest(prompt="hi") + ) + + assert excinfo.value.type == "OpenRouterHTTP429" + assert not excinfo.value.non_retryable + assert excinfo.value.next_retry_delay == timedelta(seconds=7) + + +async def test_insufficient_credits_is_non_retryable() -> None: + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 402, json={"error": {"code": 402, "message": "Insufficient credits"}} + ) + + with pytest.raises(ApplicationError) as excinfo: + await ActivityEnvironment().run( + make_activities(handler).call_openrouter, OpenRouterRequest(prompt="hi") + ) + + assert excinfo.value.type == "OpenRouterOutOfCredits" + assert excinfo.value.non_retryable + assert excinfo.value.message == "OpenRouter returned HTTP 402: Insufficient credits" + + +async def test_in_flight_budget_402_is_transient() -> None: + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 402, + json={ + "error": { + "code": 402, + "message": "In-flight budget exceeded", + "metadata": {"limit_source": "openrouter_in_flight_budget"}, + } + }, + headers={"Retry-After": "3"}, + ) + + with pytest.raises(ApplicationError) as excinfo: + await ActivityEnvironment().run( + make_activities(handler).call_openrouter, OpenRouterRequest(prompt="hi") + ) + + assert excinfo.value.type == "OpenRouterHTTP402" + assert not excinfo.value.non_retryable + assert excinfo.value.next_retry_delay == timedelta(seconds=3) + + +async def test_provider_error_on_the_choice_is_not_an_answer() -> None: + def handler(request: httpx.Request) -> httpx.Response: + body = completion_body(answer="partial ans") + body["choices"][0]["finish_reason"] = "error" + body["choices"][0]["error"] = {"code": 502, "message": "Provider died"} + return httpx.Response(200, json=body) + + with pytest.raises(ApplicationError) as excinfo: + await ActivityEnvironment().run( + make_activities(handler).call_openrouter, OpenRouterRequest(prompt="hi") + ) + + assert excinfo.value.type == "OpenRouterHTTP502" + assert not excinfo.value.non_retryable + assert "Provider died" in excinfo.value.message + + +async def test_empty_or_negative_retry_after_is_ignored() -> None: + for value in ("", "-5", "0"): + + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 429, + json={"error": {"code": 429, "message": "Rate limited"}}, + headers={"Retry-After": value}, + ) + + with pytest.raises(ApplicationError) as excinfo: + await ActivityEnvironment().run( + make_activities(handler).call_openrouter, OpenRouterRequest(prompt="hi") + ) + assert excinfo.value.next_retry_delay is None, repr(value) + + +async def test_key_limit_exceeded_is_out_of_credits_too() -> None: + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 403, + json={ + "error": {"code": 403, "message": "Key limit exceeded (total limit)"} + }, + ) + + with pytest.raises(ApplicationError) as excinfo: + await ActivityEnvironment().run( + make_activities(handler).call_openrouter, OpenRouterRequest(prompt="hi") + ) + + assert excinfo.value.type == "OpenRouterOutOfCredits" + assert excinfo.value.non_retryable + assert ( + excinfo.value.message + == "OpenRouter returned HTTP 403: Key limit exceeded (total limit)" + ) + + +async def test_server_error_is_retryable() -> None: + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response(502, json={"error": {"code": 502, "message": "down"}}) + + with pytest.raises(ApplicationError) as excinfo: + await ActivityEnvironment().run( + make_activities(handler).call_openrouter, OpenRouterRequest(prompt="hi") + ) + + assert excinfo.value.type == "OpenRouterHTTP502" + assert not excinfo.value.non_retryable + + +async def test_error_body_inside_200_is_classified_by_its_code() -> None: + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + json={"error": {"code": 403, "message": "Flagged by moderation"}}, + ) + + with pytest.raises(ApplicationError) as excinfo: + await ActivityEnvironment().run( + make_activities(handler).call_openrouter, OpenRouterRequest(prompt="hi") + ) + + assert excinfo.value.type == "OpenRouterHTTP403" + assert excinfo.value.non_retryable + + +async def test_fail_once_after_call_fails_first_attempt_only() -> None: + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 200, + json=completion_body(cost=0), + headers={"X-OpenRouter-Cache-Status": "HIT"}, + ) + + activities = make_activities(handler) + request = OpenRouterRequest(prompt="hi", fail_once_after_call=True) + + env = ActivityEnvironment() + with pytest.raises(ApplicationError) as excinfo: + await env.run(activities.call_openrouter, request) + assert excinfo.value.type == "SimulatedFailure" + assert not excinfo.value.non_retryable + + env.info = dataclasses.replace(env.info, attempt=2) + result = await env.run(activities.call_openrouter, request) + assert result.cache_status == "HIT" + assert result.cost_usd == 0.0 + + +async def test_missing_cost_is_reported_as_unknown() -> None: + def handler(request: httpx.Request) -> httpx.Response: + body = completion_body() + del body["usage"]["cost"] + return httpx.Response(200, json=body) + + result = await ActivityEnvironment().run( + make_activities(handler).call_openrouter, OpenRouterRequest(prompt="hi") + ) + assert result.cost_usd is None + assert result.cache_status == "" + + +async def test_retry_after_http_date_becomes_next_retry_delay() -> None: + from datetime import datetime, timedelta, timezone + from email.utils import format_datetime + + when = datetime.now(timezone.utc) + timedelta(seconds=30) + + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 503, + json={"error": {"code": 503, "message": "No provider available"}}, + headers={"Retry-After": format_datetime(when, usegmt=True)}, + ) + + with pytest.raises(ApplicationError) as excinfo: + await ActivityEnvironment().run( + make_activities(handler).call_openrouter, OpenRouterRequest(prompt="hi") + ) + + assert excinfo.value.type == "OpenRouterHTTP503" + assert excinfo.value.next_retry_delay is not None + assert ( + timedelta(seconds=25) < excinfo.value.next_retry_delay <= timedelta(seconds=30) + ) + + +async def test_huge_retry_after_is_capped() -> None: + def handler(request: httpx.Request) -> httpx.Response: + return httpx.Response( + 429, + json={"error": {"code": 429, "message": "Rate limited"}}, + headers={"Retry-After": "1000000000"}, + ) + + with pytest.raises(ApplicationError) as excinfo: + await ActivityEnvironment().run( + make_activities(handler).call_openrouter, OpenRouterRequest(prompt="hi") + ) + assert excinfo.value.next_retry_delay == timedelta(minutes=5) + + +async def test_no_choices_and_no_error_is_retried() -> None: + def handler(request: httpx.Request) -> httpx.Response: + body = completion_body() + body["choices"] = [] + return httpx.Response(200, json=body) + + with pytest.raises(ApplicationError) as excinfo: + await ActivityEnvironment().run( + make_activities(handler).call_openrouter, OpenRouterRequest(prompt="hi") + ) + assert excinfo.value.type == "OpenRouterHTTP500" + assert not excinfo.value.non_retryable diff --git a/tests/openrouter/budget_gate_test.py b/tests/openrouter/budget_gate_test.py new file mode 100644 index 000000000..faa16cc34 --- /dev/null +++ b/tests/openrouter/budget_gate_test.py @@ -0,0 +1,610 @@ +import asyncio +import uuid + +import pytest +from temporalio import activity +from temporalio.client import ( + Client, + WithStartWorkflowOperation, + WorkflowFailureError, + WorkflowHandle, + WorkflowUpdateFailedError, +) +from temporalio.common import WorkflowIDConflictPolicy +from temporalio.exceptions import ApplicationError, CancelledError +from temporalio.worker import Worker + +from openrouter.budget_gate.workflow import BudgetGateWorkflow +from openrouter.shared import ( + BatchInput, + BatchResult, + BudgetGateInput, + OpenRouterRequest, + OpenRouterResult, + SpendReport, +) + +COST_PER_CALL = 0.001 + + +class FakeOpenRouter: + """Mock Activity with a per-prompt count; optionally out of credits on first call.""" + + def __init__(self, out_of_credits_for: set[str] | None = None) -> None: + self.calls: dict[str, int] = {} + self.out_of_credits_for = out_of_credits_for or set() + + @activity.defn(name="call_openrouter") + async def call_openrouter(self, request: OpenRouterRequest) -> OpenRouterResult: + self.calls[request.prompt] = self.calls.get(request.prompt, 0) + 1 + if ( + request.prompt in self.out_of_credits_for + and self.calls[request.prompt] == 1 + ): + raise ApplicationError( + "OpenRouter returned HTTP 403: Key limit exceeded (total limit)", + type="OpenRouterOutOfCredits", + non_retryable=True, + ) + return OpenRouterResult( + prompt=request.prompt, + model="openai/gpt-4o-mini", + answer="ok", + cost_usd=COST_PER_CALL, + generation_id=f"gen-{request.prompt}-{self.calls[request.prompt]}", + cache_status="MISS", + ) + + +@pytest.fixture +def task_queue() -> str: + return f"test-openrouter-budget-{uuid.uuid4()}" + + +async def start( + client: Client, task_queue: str, gate: BudgetGateInput +) -> WorkflowHandle[BudgetGateWorkflow, BatchResult]: + return await client.start_workflow( + BudgetGateWorkflow.run, + gate, + id=f"test-openrouter-budget-{uuid.uuid4()}", + task_queue=task_queue, + ) + + +async def wait_until_paused( + handle: WorkflowHandle[BudgetGateWorkflow, BatchResult], reason: str +) -> SpendReport: + for _ in range(100): + report = await handle.query(BudgetGateWorkflow.spend_report) + if reason in report.paused.values(): + return report + await asyncio.sleep(0.1) + raise AssertionError(f"workflow never paused with reason {reason!r}") + + +async def test_soft_budget_pauses_then_resumes_on_raise_budget( + client: Client, task_queue: str +) -> None: + fake = FakeOpenRouter() + async with Worker( + client, + task_queue=task_queue, + workflows=[BudgetGateWorkflow], + activities=[fake.call_openrouter], + ): + # Budget covers exactly one call; the second prompt must park. + handle = await start( + client, + task_queue, + BudgetGateInput( + batch=BatchInput(prompts=["a", "b", "c"], max_concurrency=1), + budget_usd=0.0015, + estimated_cost_usd=COST_PER_CALL, + approval_timeout_seconds=60, + ), + ) + report = await wait_until_paused(handle, "soft_budget_exhausted") + assert report.completed == 1 + assert report.spent_usd == pytest.approx(COST_PER_CALL) + assert report.paused == {"b": "soft_budget_exhausted"} + + report = await handle.execute_update(BudgetGateWorkflow.raise_budget, 0.01) + assert report.budget_usd == 0.01 + + result = await handle.result() + + assert [r.prompt for r in result.results] == ["a", "b", "c"] + assert result.skipped == [] + assert result.reported_cost_usd == pytest.approx(3 * COST_PER_CALL) + assert fake.calls == {"a": 1, "b": 1, "c": 1} + + +async def test_insufficient_credits_pauses_and_reruns_same_prompt( + client: Client, task_queue: str +) -> None: + fake = FakeOpenRouter(out_of_credits_for={"b"}) + async with Worker( + client, + task_queue=task_queue, + workflows=[BudgetGateWorkflow], + activities=[fake.call_openrouter], + ): + handle = await start( + client, + task_queue, + BudgetGateInput( + batch=BatchInput(prompts=["a", "b", "c"], max_concurrency=1), + budget_usd=1.0, + estimated_cost_usd=COST_PER_CALL, + approval_timeout_seconds=60, + ), + ) + report = await wait_until_paused(handle, "insufficient_credits") + assert report.paused == {"b": "insufficient_credits"} + assert report.completed == 1 + + # Re-sending the same budget is how an operator says "I topped up". + await handle.execute_update(BudgetGateWorkflow.raise_budget, 1.0) + result = await handle.result() + + assert [r.prompt for r in result.results] == ["a", "b", "c"] + assert result.skipped == [] + # "b" was called twice: once for the 402, once after the budget bump. + assert fake.calls == {"a": 1, "b": 2, "c": 1} + assert result.results[1].generation_id == "gen-b-2" + + +async def test_lowering_the_budget_is_rejected(client: Client, task_queue: str) -> None: + fake = FakeOpenRouter() + async with Worker( + client, + task_queue=task_queue, + workflows=[BudgetGateWorkflow], + activities=[fake.call_openrouter], + ): + handle = await start( + client, + task_queue, + BudgetGateInput( + batch=BatchInput(prompts=["a", "b"], max_concurrency=1), + budget_usd=0.0015, + estimated_cost_usd=COST_PER_CALL, + approval_timeout_seconds=60, + ), + ) + await wait_until_paused(handle, "soft_budget_exhausted") + + with pytest.raises(WorkflowUpdateFailedError): + await handle.execute_update(BudgetGateWorkflow.raise_budget, 0.0001) + + await handle.execute_update(BudgetGateWorkflow.raise_budget, 0.01) + result = await handle.result() + + assert [r.prompt for r in result.results] == ["a", "b"] + + +async def test_approval_timeout_skips_remaining_prompts( + client: Client, task_queue: str +) -> None: + fake = FakeOpenRouter() + async with Worker( + client, + task_queue=task_queue, + workflows=[BudgetGateWorkflow], + activities=[fake.call_openrouter], + ): + result = await client.execute_workflow( + BudgetGateWorkflow.run, + BudgetGateInput( + batch=BatchInput(prompts=["a", "b", "c"], max_concurrency=2), + budget_usd=0.0015, + estimated_cost_usd=COST_PER_CALL, + approval_timeout_seconds=1, + ), + id=f"test-openrouter-budget-{uuid.uuid4()}", + task_queue=task_queue, + ) + + assert [r.prompt for r in result.results] == ["a"] + assert sorted(s.prompt for s in result.skipped) == ["b", "c"] + assert {s.reason for s in result.skipped} == {"soft_budget_exhausted"} + assert result.reported_cost_usd == pytest.approx(COST_PER_CALL) + assert fake.calls == {"a": 1} + + +async def test_zero_concurrency_is_rejected(client: Client, task_queue: str) -> None: + fake = FakeOpenRouter() + async with Worker( + client, + task_queue=task_queue, + workflows=[BudgetGateWorkflow], + activities=[fake.call_openrouter], + ): + with pytest.raises(WorkflowFailureError) as excinfo: + await client.execute_workflow( + BudgetGateWorkflow.run, + BudgetGateInput( + batch=BatchInput(prompts=["a"], max_concurrency=0), + budget_usd=1.0, + ), + id=f"test-openrouter-budget-{uuid.uuid4()}", + task_queue=task_queue, + ) + assert isinstance(excinfo.value.cause, ApplicationError) + assert "max_concurrency" in str(excinfo.value.cause) + assert fake.calls == {} + + +async def test_non_finite_budget_is_rejected(client: Client, task_queue: str) -> None: + fake = FakeOpenRouter() + async with Worker( + client, + task_queue=task_queue, + workflows=[BudgetGateWorkflow], + activities=[fake.call_openrouter], + ): + with pytest.raises(WorkflowFailureError): + await client.execute_workflow( + BudgetGateWorkflow.run, + BudgetGateInput( + batch=BatchInput(prompts=["a"]), budget_usd=float("inf") + ), + id=f"test-openrouter-budget-{uuid.uuid4()}", + task_queue=task_queue, + ) + assert fake.calls == {} + + +async def test_raising_the_budget_admits_only_what_fits( + client: Client, task_queue: str +) -> None: + """Three prompts park; a raise that fits one call must wake only one.""" + fake = FakeOpenRouter() + async with Worker( + client, + task_queue=task_queue, + workflows=[BudgetGateWorkflow], + activities=[fake.call_openrouter], + ): + handle = await start( + client, + task_queue, + BudgetGateInput( + batch=BatchInput(prompts=["a", "b", "c", "d"], max_concurrency=3), + budget_usd=0.0015, + estimated_cost_usd=COST_PER_CALL, + approval_timeout_seconds=60, + ), + ) + # "a" runs; "b" and "c" park on the soft budget, and "d" parks too + # once "a" frees its concurrency slot. + for _ in range(100): + report = await handle.query(BudgetGateWorkflow.spend_report) + if report.completed == 1 and len(report.paused) == 3: + break + await asyncio.sleep(0.1) + assert report.completed == 1 and len(report.paused) == 3 + + # Headroom for exactly one more call: spent 0.001 + 0.001 <= 0.0025. + await handle.execute_update(BudgetGateWorkflow.raise_budget, 0.0025) + for _ in range(100): + report = await handle.query(BudgetGateWorkflow.spend_report) + if report.completed == 2: + break + await asyncio.sleep(0.1) + # Give the others a chance to (wrongly) run; they must stay parked. + await asyncio.sleep(0.5) + report = await handle.query(BudgetGateWorkflow.spend_report) + assert report.completed == 2 + assert report.spent_usd == pytest.approx(2 * COST_PER_CALL) + assert len(report.paused) == 2 + assert report.reserved_usd == 0 + + await handle.execute_update(BudgetGateWorkflow.raise_budget, 1.0) + result = await handle.result() + + assert [r.prompt for r in result.results] == ["a", "b", "c", "d"] + assert fake.calls == {"a": 1, "b": 1, "c": 1, "d": 1} + + +async def test_update_with_start_sets_the_budget_before_run( + client: Client, task_queue: str +) -> None: + """An Update that lands before run() must not be overwritten by run().""" + fake = FakeOpenRouter() + async with Worker( + client, + task_queue=task_queue, + workflows=[BudgetGateWorkflow], + activities=[fake.call_openrouter], + ): + start_op = WithStartWorkflowOperation( + BudgetGateWorkflow.run, + BudgetGateInput( + batch=BatchInput(prompts=["a", "b", "c"], max_concurrency=1), + budget_usd=0.0015, + estimated_cost_usd=COST_PER_CALL, + approval_timeout_seconds=5, + ), + id=f"test-openrouter-budget-{uuid.uuid4()}", + task_queue=task_queue, + id_conflict_policy=WorkflowIDConflictPolicy.FAIL, + ) + report = await client.execute_update_with_start_workflow( + BudgetGateWorkflow.raise_budget, 1.0, start_workflow_operation=start_op + ) + assert report.budget_usd == 1.0 + handle = await start_op.workflow_handle() + result = await handle.result() + + # With the raised budget honored, nothing parks and nothing is skipped. + assert [r.prompt for r in result.results] == ["a", "b", "c"] + assert result.skipped == [] + + +async def test_cancellation_is_not_a_skipped_prompt( + client: Client, task_queue: str, caplog: pytest.LogCaptureFixture +) -> None: + @activity.defn(name="call_openrouter") + async def slow_call(request: OpenRouterRequest) -> OpenRouterResult: + # Heartbeat so the cancellation request reaches the Activity. + while True: + activity.heartbeat() + await asyncio.sleep(0.1) + + async with Worker( + client, + task_queue=task_queue, + workflows=[BudgetGateWorkflow], + activities=[slow_call], + ): + handle = await start( + client, + task_queue, + BudgetGateInput( + batch=BatchInput(prompts=["a", "b"], max_concurrency=2), + budget_usd=1.0, + estimated_cost_usd=COST_PER_CALL, + approval_timeout_seconds=60, + ), + ) + await asyncio.sleep(0.5) + await handle.cancel() + with pytest.raises(WorkflowFailureError) as excinfo: + await handle.result() + + assert isinstance(excinfo.value.cause, CancelledError) + # The cancelled Activities must not have been recorded as skipped prompts. + assert not [r for r in caplog.records if "Skipping prompt" in r.getMessage()] + + +async def test_bad_approval_timeout_fails_the_workflow( + client: Client, task_queue: str +) -> None: + fake = FakeOpenRouter() + async with Worker( + client, + task_queue=task_queue, + workflows=[BudgetGateWorkflow], + activities=[fake.call_openrouter], + ): + for bad in (-5, 10**15): + with pytest.raises(WorkflowFailureError) as excinfo: + await client.execute_workflow( + BudgetGateWorkflow.run, + BudgetGateInput( + batch=BatchInput(prompts=["a"]), + budget_usd=1.0, + approval_timeout_seconds=bad, + ), + id=f"test-openrouter-budget-{uuid.uuid4()}", + task_queue=task_queue, + ) + assert isinstance(excinfo.value.cause, ApplicationError) + assert "approval_timeout_seconds" in str(excinfo.value.cause) + assert fake.calls == {} + + +async def test_unknown_cost_is_charged_at_the_estimate( + client: Client, task_queue: str +) -> None: + @activity.defn(name="call_openrouter") + async def costless(request: OpenRouterRequest) -> OpenRouterResult: + return OpenRouterResult( + prompt=request.prompt, + model="m", + answer="ok", + cost_usd=None, + generation_id="gen", + cache_status="", + ) + + async with Worker( + client, + task_queue=task_queue, + workflows=[BudgetGateWorkflow], + activities=[costless], + ): + result = await client.execute_workflow( + BudgetGateWorkflow.run, + BudgetGateInput( + batch=BatchInput(prompts=["a", "b"], max_concurrency=1), + budget_usd=1.0, + estimated_cost_usd=0.002, + ), + id=f"test-openrouter-budget-{uuid.uuid4()}", + task_queue=task_queue, + ) + + assert result.reported_cost_usd == pytest.approx(0.004) + assert result.unknown_cost_count == 2 + assert [r.cost_usd for r in result.results] == [None, None] + + +async def test_credits_never_arrive_skips_with_reason( + client: Client, task_queue: str +) -> None: + fake = FakeOpenRouter(out_of_credits_for={"a"}) + async with Worker( + client, + task_queue=task_queue, + workflows=[BudgetGateWorkflow], + activities=[fake.call_openrouter], + ): + result = await client.execute_workflow( + BudgetGateWorkflow.run, + BudgetGateInput( + batch=BatchInput(prompts=["a"], max_concurrency=1), + budget_usd=1.0, + estimated_cost_usd=COST_PER_CALL, + approval_timeout_seconds=1, + ), + id=f"test-openrouter-budget-{uuid.uuid4()}", + task_queue=task_queue, + ) + + assert result.results == [] + assert [(s.prompt, s.reason) for s in result.skipped] == [ + ("a", "insufficient_credits") + ] + + +async def test_exact_budget_boundary_is_affordable( + client: Client, task_queue: str +) -> None: + """Ten $0.001 calls must fit a $0.01 budget despite float accumulation.""" + fake = FakeOpenRouter() + async with Worker( + client, + task_queue=task_queue, + workflows=[BudgetGateWorkflow], + activities=[fake.call_openrouter], + ): + result = await client.execute_workflow( + BudgetGateWorkflow.run, + BudgetGateInput( + batch=BatchInput( + prompts=[str(i) for i in range(10)], max_concurrency=1 + ), + budget_usd=0.01, + estimated_cost_usd=COST_PER_CALL, + approval_timeout_seconds=0, + ), + id=f"test-openrouter-budget-{uuid.uuid4()}", + task_queue=task_queue, + ) + + assert len(result.results) == 10 + assert result.skipped == [] + + +async def test_top_up_during_an_in_flight_call_counts( + client: Client, task_queue: str +) -> None: + """A raise_budget that lands while a call is in flight must release that + prompt when the call then fails for lack of credits.""" + release_b = asyncio.Event() + calls: dict[str, int] = {} + + @activity.defn(name="call_openrouter") + async def flaky_credits(request: OpenRouterRequest) -> OpenRouterResult: + calls[request.prompt] = calls.get(request.prompt, 0) + 1 + if calls[request.prompt] == 1: + if request.prompt == "b": + # Hold b's failure until the test has sent the top-up. + await release_b.wait() + raise ApplicationError( + "OpenRouter returned HTTP 402: Insufficient credits", + type="OpenRouterOutOfCredits", + non_retryable=True, + ) + return OpenRouterResult( + prompt=request.prompt, + model="m", + answer="ok", + cost_usd=COST_PER_CALL, + generation_id=f"gen-{request.prompt}", + cache_status="", + ) + + async with Worker( + client, + task_queue=task_queue, + workflows=[BudgetGateWorkflow], + activities=[flaky_credits], + ): + handle = await start( + client, + task_queue, + BudgetGateInput( + batch=BatchInput(prompts=["a", "b"], max_concurrency=2), + budget_usd=1.0, + estimated_cost_usd=COST_PER_CALL, + approval_timeout_seconds=5, + ), + ) + report = await wait_until_paused(handle, "insufficient_credits") + assert report.paused == {"a": "insufficient_credits"} + # The operator tops up while b's call is still in flight... + await handle.execute_update(BudgetGateWorkflow.raise_budget, 1.0) + # ...and only then does b's failure reach the Workflow. + release_b.set() + result = await handle.result() + + assert [r.prompt for r in result.results] == ["a", "b"] + assert result.skipped == [] + assert calls == {"a": 2, "b": 2} + + +async def test_top_up_during_a_call_counts_even_after_the_deadline( + client: Client, task_queue: str +) -> None: + release = asyncio.Event() + calls: dict[str, int] = {} + + @activity.defn(name="call_openrouter") + async def flaky_credits(request: OpenRouterRequest) -> OpenRouterResult: + calls[request.prompt] = calls.get(request.prompt, 0) + 1 + if calls[request.prompt] == 1: + await release.wait() + raise ApplicationError( + "OpenRouter returned HTTP 402: Insufficient credits", + type="OpenRouterOutOfCredits", + non_retryable=True, + ) + return OpenRouterResult( + prompt=request.prompt, + model="m", + answer="ok", + cost_usd=COST_PER_CALL, + generation_id="gen", + cache_status="", + ) + + async with Worker( + client, + task_queue=task_queue, + workflows=[BudgetGateWorkflow], + activities=[flaky_credits], + ): + handle = await start( + client, + task_queue, + BudgetGateInput( + batch=BatchInput(prompts=["a"]), + budget_usd=1.0, + estimated_cost_usd=COST_PER_CALL, + approval_timeout_seconds=0, + ), + ) + for _ in range(100): + if calls.get("a") == 1: + break + await asyncio.sleep(0.05) + await handle.execute_update(BudgetGateWorkflow.raise_budget, 1.0) + release.set() + result = await handle.result() + + assert [r.prompt for r in result.results] == ["a"] + assert calls == {"a": 2} diff --git a/tests/openrouter/prompt_batch_test.py b/tests/openrouter/prompt_batch_test.py new file mode 100644 index 000000000..a744cfd18 --- /dev/null +++ b/tests/openrouter/prompt_batch_test.py @@ -0,0 +1,95 @@ +import asyncio +import uuid + +import pytest +from temporalio import activity +from temporalio.client import Client, WorkflowFailureError +from temporalio.exceptions import ApplicationError, CancelledError +from temporalio.worker import Worker + +from openrouter.prompt_batch.workflow import PromptBatchWorkflow +from openrouter.shared import BatchInput, OpenRouterRequest, OpenRouterResult + + +def fake_result(request: OpenRouterRequest, cost: float = 0.001) -> OpenRouterResult: + return OpenRouterResult( + prompt=request.prompt, + model="openai/gpt-4o-mini", + answer=f"Answer to: {request.prompt}", + cost_usd=cost, + generation_id=f"gen-{request.prompt}", + cache_status="MISS", + ) + + +async def test_prompt_batch_collects_results_and_skips_failures( + client: Client, +) -> None: + @activity.defn(name="call_openrouter") + async def mock_call_openrouter(request: OpenRouterRequest) -> OpenRouterResult: + if request.prompt == "bad": + raise ApplicationError( + "OpenRouter returned HTTP 400: bad request", + type="OpenRouterHTTP400", + non_retryable=True, + ) + if request.prompt == "costless": + result = fake_result(request) + result.cost_usd = None + return result + return fake_result(request) + + task_queue = f"test-openrouter-{uuid.uuid4()}" + async with Worker( + client, + task_queue=task_queue, + workflows=[PromptBatchWorkflow], + activities=[mock_call_openrouter], + ): + result = await client.execute_workflow( + PromptBatchWorkflow.run, + BatchInput(prompts=["one", "bad", "two", "costless"], max_concurrency=2), + id=f"test-openrouter-{uuid.uuid4()}", + task_queue=task_queue, + ) + + assert [r.prompt for r in result.results] == ["one", "two", "costless"] + assert all(r.answer.startswith("Answer to:") for r in result.results) + assert [(s.prompt, s.reason) for s in result.skipped] == [ + ("bad", "OpenRouterHTTP400") + ] + # Subtotal of the known costs, with the unknown one counted separately. + assert result.reported_cost_usd == 0.002 + assert result.unknown_cost_count == 1 + + +async def test_cancellation_is_not_a_skipped_prompt( + client: Client, caplog: pytest.LogCaptureFixture +) -> None: + @activity.defn(name="call_openrouter") + async def slow_call(request: OpenRouterRequest) -> OpenRouterResult: + while True: + activity.heartbeat() + await asyncio.sleep(0.1) + + task_queue = f"test-openrouter-{uuid.uuid4()}" + async with Worker( + client, + task_queue=task_queue, + workflows=[PromptBatchWorkflow], + activities=[slow_call], + ): + handle = await client.start_workflow( + PromptBatchWorkflow.run, + BatchInput(prompts=["one", "two"], max_concurrency=2), + id=f"test-openrouter-{uuid.uuid4()}", + task_queue=task_queue, + ) + await asyncio.sleep(0.5) + await handle.cancel() + with pytest.raises(WorkflowFailureError) as excinfo: + await handle.result() + + assert isinstance(excinfo.value.cause, CancelledError) + # The cancelled Activities must not have been recorded as skipped prompts. + assert not [r for r in caplog.records if "Skipping prompt" in r.getMessage()] diff --git a/uv.lock b/uv.lock index 36f5bbb6c..218d820a5 100644 --- a/uv.lock +++ b/uv.lock @@ -4658,6 +4658,9 @@ openai-agents = [ { name = "requests" }, { name = "temporalio-openai-agents" }, ] +openrouter = [ + { name = "openai" }, +] pydantic-converter = [ { name = "pydantic" }, ] @@ -4762,6 +4765,7 @@ openai-agents = [ { name = "requests", specifier = ">=2.32.0,<3" }, { name = "temporalio-openai-agents", specifier = ">=1.0.0" }, ] +openrouter = [{ name = "openai", specifier = ">=1.4.0,<3" }] pydantic-converter = [{ name = "pydantic", specifier = ">=2.10.6,<3" }] sentry = [{ name = "sentry-sdk", specifier = ">=2.13.0" }] strands-agents = [