diff --git a/.github/codeql/codeql-config.yml b/.github/codeql/codeql-config.yml index 6e15c1069a3..61982be90a9 100644 --- a/.github/codeql/codeql-config.yml +++ b/.github/codeql/codeql-config.yml @@ -1,8 +1,5 @@ name: "LiteLLM CodeQL config" -queries: - - uses: security-and-quality - # Known OOM queries on large Python codebases: # CodeQL builds a full data flow graph in memory. These two queries trace # sensitive data through every log call / regex pattern, causing combinatorial @@ -14,17 +11,6 @@ query-filters: id: py/clear-text-logging-sensitive-data # CWE-312 - exclude: id: py/polynomial-redos # CWE-730 - # Import resolution confuses stdlib types with management_endpoints/types.py. - # The generic cycle query also reports intentional deferred imports. - - exclude: - id: py/cyclic-import - - exclude: - id: py/unsafe-cyclic-import - # Known false positives on live settings and Protocol placeholders. - - exclude: - id: py/unused-global-variable - - exclude: - id: py/ineffectual-statement paths-ignore: - tests diff --git a/.github/workflows/codeql.yml b/.github/workflows/codeql.yml index 9a85ced57f6..7fd767abf59 100644 --- a/.github/workflows/codeql.yml +++ b/.github/workflows/codeql.yml @@ -1,8 +1,6 @@ name: "CodeQL" on: - push: - branches: [main] pull_request: branches: [main] schedule: @@ -43,14 +41,15 @@ jobs: persist-credentials: false - name: Initialize CodeQL - uses: github/codeql-action/init@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3.34.1 + uses: github/codeql-action/init@2892aa5e19bbd11bc0cff5427e3b750a04d9e3c2 # v4.38.2 with: languages: ${{ matrix.language }} build-mode: ${{ matrix.build-mode }} config-file: ./.github/codeql/codeql-config.yml + queries: ${{ github.event_name == 'pull_request' && '+security-extended' || '' }} - name: Perform CodeQL Analysis - uses: github/codeql-action/analyze@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3.34.1 + uses: github/codeql-action/analyze@2892aa5e19bbd11bc0cff5427e3b750a04d9e3c2 # v4.38.2 with: category: "/language:${{ matrix.language }}" output: sarif-results @@ -83,7 +82,7 @@ jobs: output: sarif-results/python.sarif - name: Upload SARIF - uses: github/codeql-action/upload-sarif@ebcb5b36ded6beda4ceefea6a8bc4cc885255bb3 # v3.34.1 + uses: github/codeql-action/upload-sarif@2892aa5e19bbd11bc0cff5427e3b750a04d9e3c2 # v4.38.2 with: sarif_file: sarif-results category: "/language:${{ matrix.language }}" diff --git a/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000200_lens_trace_signals/migration.sql b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000200_lens_trace_signals/migration.sql new file mode 100644 index 00000000000..67043eccb5f --- /dev/null +++ b/litellm-proxy-extras/litellm_proxy_extras/migrations/20261008000200_lens_trace_signals/migration.sql @@ -0,0 +1,16 @@ +CREATE TABLE IF NOT EXISTS "LiteLLM_LensSignalConfig" ( + "id" TEXT NOT NULL, + "data" JSONB NOT NULL, + CONSTRAINT "LiteLLM_LensSignalConfig_pkey" PRIMARY KEY ("id") +); + +CREATE TABLE IF NOT EXISTS "LiteLLM_LensTraceSignal" ( + "trace_id" TEXT NOT NULL, + "trace_ref" TEXT NOT NULL DEFAULT '', + "config_key" TEXT NOT NULL, + "span_count" INTEGER NOT NULL, + "claimed_until" TIMESTAMP(3), + "classified_at" TIMESTAMP(3), + "data" JSONB NOT NULL, + CONSTRAINT "LiteLLM_LensTraceSignal_pkey" PRIMARY KEY ("trace_id", "trace_ref") +); diff --git a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma index 2f1d4aab78e..1b5d870ce1b 100644 --- a/litellm-proxy-extras/litellm_proxy_extras/schema.prisma +++ b/litellm-proxy-extras/litellm_proxy_extras/schema.prisma @@ -1982,3 +1982,20 @@ model LiteLLM_LensDataset { @@id([id, revision]) } + +model LiteLLM_LensSignalConfig { + id String @id + data Json +} + +model LiteLLM_LensTraceSignal { + trace_id String + trace_ref String @default("") + config_key String + span_count Int + claimed_until DateTime? + classified_at DateTime? + data Json + + @@id([trace_id, trace_ref]) +} diff --git a/litellm-rust/crates/traces-clickhouse/query/lens_content.sql b/litellm-rust/crates/traces-clickhouse/query/lens_content.sql index 99fb56a5f48..52355a11061 100644 --- a/litellm-rust/crates/traces-clickhouse/query/lens_content.sql +++ b/litellm-rust/crates/traces-clickhouse/query/lens_content.sql @@ -17,6 +17,7 @@ SELECT * FROM ( FROM otel_traces WHERE {source:String}='traces' AND ({all_teams:UInt8}=1 OR TeamId={team:String}) AND ({key_hash:String}='' OR ApiKeyHash={key_hash:String}) + AND Timestamp >= parseDateTime64BestEffortOrZero({start_time:String}, 9) - INTERVAL 7 DAY AND ({trace_ref:String}='' OR hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId)))={trace_ref:String}) AND TraceId={id:String} AND TeamId={record_team:String} AND SpanId > {cursor:String} ORDER BY SpanId LIMIT 1 BY SpanId LIMIT 40 @@ -35,5 +36,6 @@ SELECT * FROM ( FROM spend_logs FINAL WHERE {source:String}='requests' AND ({all_teams:UInt8}=1 OR team_id={team:String}) AND ({key_hash:String}='' OR api_key={key_hash:String}) + AND spend_logs.start_time >= parseDateTime64BestEffortOrZero({start_time:String}, 3) - INTERVAL 7 DAY AND request_id={id:String} AND team_id={record_team:String} LIMIT 1 ) diff --git a/litellm-rust/crates/traces-clickhouse/query/lens_evidence.sql b/litellm-rust/crates/traces-clickhouse/query/lens_evidence.sql index a0d600cdfde..b53617364cf 100644 --- a/litellm-rust/crates/traces-clickhouse/query/lens_evidence.sql +++ b/litellm-rust/crates/traces-clickhouse/query/lens_evidence.sql @@ -2,6 +2,7 @@ SELECT sum(matches) AS count FROM ( SELECT count() AS matches FROM otel_traces WHERE {source:String}='traces' AND ({all_teams:UInt8}=1 OR TeamId={team:String}) AND ({key_hash:String}='' OR ApiKeyHash={key_hash:String}) + AND Timestamp >= parseDateTime64BestEffortOrZero({start_time:String}, 9) - INTERVAL 7 DAY AND ({trace_ref:String}='' OR hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId)))={trace_ref:String}) AND TraceId={id:String} AND TeamId={record_team:String} AND SpanId={span:String} AND position(concat('Input: ',Input,'\nOutput: ',Output,'\nStatus: ',StatusCode,' ',StatusMessage),{quote:String})>0 @@ -9,6 +10,7 @@ SELECT sum(matches) AS count FROM ( SELECT count() AS matches FROM spend_logs FINAL WHERE {source:String}='requests' AND ({all_teams:UInt8}=1 OR team_id={team:String}) AND ({key_hash:String}='' OR api_key={key_hash:String}) + AND spend_logs.start_time >= parseDateTime64BestEffortOrZero({start_time:String}, 3) - INTERVAL 7 DAY AND request_id={id:String} AND team_id={record_team:String} AND request_id={span:String} AND position(concat('Input: ',messages,'\nOutput: ',response,'\nError: ',error_str),{quote:String})>0 ) diff --git a/litellm-rust/crates/traces-clickhouse/src/query/lens.rs b/litellm-rust/crates/traces-clickhouse/src/query/lens.rs index 6439e696dad..77474c7143d 100644 --- a/litellm-rust/crates/traces-clickhouse/src/query/lens.rs +++ b/litellm-rust/crates/traces-clickhouse/src/query/lens.rs @@ -209,6 +209,7 @@ pub struct LensContentParams { pub source: ContentSource, pub id: String, pub record_team: String, + pub start_time: String, pub trace_ref: String, pub cursor: String, #[serde(deserialize_with = "super::number::deserialize")] @@ -252,6 +253,7 @@ pub struct LensEvidenceParams { pub source: ContentSource, pub id: String, pub record_team: String, + pub start_time: String, pub trace_ref: String, pub span: String, pub quote: String, diff --git a/litellm-rust/crates/traces-clickhouse/src/query/number.rs b/litellm-rust/crates/traces-clickhouse/src/query/number.rs index 9283903fee1..7a07845d0cd 100644 --- a/litellm-rust/crates/traces-clickhouse/src/query/number.rs +++ b/litellm-rust/crates/traces-clickhouse/src/query/number.rs @@ -99,7 +99,7 @@ mod tests { fn content_rejects_unsupported_sources(#[case] source: &str, #[case] valid: bool) { let parameters = serde_json::json!({ "all_teams": 0, "team": "team", "key_hash": "", "source": source, "id": "id", - "record_team": "team", "trace_ref": "", "cursor": "", "offset": 0 + "record_team": "team", "start_time": "", "trace_ref": "", "cursor": "", "offset": 0 }); assert_eq!( serde_json::from_value::(parameters).is_ok(), diff --git a/litellm-rust/crates/traces-clickhouse/tests/load.rs b/litellm-rust/crates/traces-clickhouse/tests/load.rs index 07c1095dfc3..aaec2c17e54 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/load.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/load.rs @@ -25,7 +25,8 @@ async fn seed_days(fixture: &SeededDatabase, first_day: u64, days: u64) -> TestR "INSERT INTO {DATABASE}.otel_traces \ (Timestamp, TraceId, SpanId, ParentSpanId, SpanName, ServiceName, ObservationType, TeamId, ApiKeyHash, Duration, SpanAttributes) \ SELECT now64(9) - toIntervalHour(intDiv(number, {SPANS_PER_DAY}) * 24 + 12 + {first_day} * 24), \ - concat('load-', toString(number + {first_row})), concat('span-', toString(number + {first_row})), \ + if({first_day} = 0, concat('load-', toString(number + {first_row})), 'load-0'), \ + concat('span-', toString(number + {first_row})), \ '', 'span', 'service', 'agent', 'load-team', '', 0, \ if({first_day}=0 AND number < {SPANS_PER_DAY}, map('payload', repeat('x', 3000)), map()) \ FROM numbers({count})" @@ -41,6 +42,62 @@ async fn seed_days(fixture: &SeededDatabase, first_day: u64, days: u64) -> TestR Ok(()) } +async fn trace_start_time(fixture: &SeededDatabase) -> TestResult { + let query = format!( + "SELECT toString(Timestamp, 'UTC') AS start_time FROM {DATABASE}.otel_traces \ + WHERE TraceId = 'load-0' LIMIT 1 FORMAT JSON" + ); + let response = fixture + .database + .client + .post(&fixture.database.url) + .body(query) + .send() + .await? + .error_for_status()? + .text() + .await?; + let result: Value = serde_json::from_str(&response)?; + result["data"][0]["start_time"] + .as_str() + .map(str::to_owned) + .ok_or_else(|| "trace start time missing".into()) +} + +fn content_parameters(start_time: &str) -> BTreeMap { + BTreeMap::from([ + ("source".into(), Parameter::Text("traces".into())), + ("all_teams".into(), Parameter::Integer(0)), + ("team".into(), Parameter::Text("load-team".into())), + ("key_hash".into(), Parameter::Text(String::new())), + ("id".into(), Parameter::Text("load-0".into())), + ("record_team".into(), Parameter::Text("load-team".into())), + ("start_time".into(), Parameter::Text(start_time.into())), + ("trace_ref".into(), Parameter::Text(String::new())), + ("cursor".into(), Parameter::Text(String::new())), + ("offset".into(), Parameter::Integer(1)), + ]) +} + +async fn content(fixture: &SeededDatabase, start_time: &str, query_id: &str) -> TestResult { + let connection = Connection::configured( + &format!("{}?query_id={query_id}", fixture.database.url), + DATABASE, + "default", + "", + )?; + let response = execute_named_read( + &fixture.database.client, + &connection, + ReadQuery::Content, + &content_parameters(start_time), + ) + .await?; + let result: Value = serde_json::from_str(&response)?; + assert!(!result["data"].as_array().ok_or("content rows")?.is_empty()); + Ok(()) +} + fn sample_parameters(start: u64, end: u64) -> BTreeMap { BTreeMap::from([ ("source".into(), Parameter::Text("traces".into())), @@ -144,3 +201,27 @@ async fn lens_sample_reads_scale_with_window_not_retention( ); Ok(()) } + +#[rstest] +#[tokio::test] +async fn lens_content_reads_scale_with_trace_not_retention( + #[future(awt)] migrated_database: TestResult, +) -> TestResult { + let fixture = migrated_database?; + seed_days(&fixture, 0, 8).await?; + let start_time = trace_start_time(&fixture).await?; + let before_id = format!("lens_content_before_{}", std::process::id()); + content(&fixture, &start_time, &before_id).await?; + let before = query_read_rows(&fixture, &before_id).await?; + + seed_days(&fixture, 8, 24).await?; + let after_id = format!("lens_content_after_{}", std::process::id()); + content(&fixture, &start_time, &after_id).await?; + let after = query_read_rows(&fixture, &after_id).await?; + println!("lens_content read_rows: before={before}, after={after}"); + assert!( + after * 100 <= before * 105, + "read_rows grew from {before} to {after}" + ); + Ok(()) +} diff --git a/litellm-rust/crates/traces-clickhouse/tests/migrations.rs b/litellm-rust/crates/traces-clickhouse/tests/migrations.rs index e51d9083c59..7630b3033a7 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/migrations.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/migrations.rs @@ -1150,6 +1150,7 @@ async fn lens_filters_reads_and_evidence_keep_reused_trace_ids_separate( ("source".into(), Parameter::Text("traces".into())), ("id".into(), Parameter::Text("shared".into())), ("record_team".into(), Parameter::Text("team".into())), + ("start_time".into(), Parameter::Text(String::new())), ("trace_ref".into(), Parameter::Text(first_ref.into())), ("cursor".into(), Parameter::Text(String::new())), ("offset".into(), Parameter::Integer(1)), @@ -1177,6 +1178,7 @@ async fn lens_filters_reads_and_evidence_keep_reused_trace_ids_separate( ("source".into(), Parameter::Text("traces".into())), ("id".into(), Parameter::Text("shared".into())), ("record_team".into(), Parameter::Text("team".into())), + ("start_time".into(), Parameter::Text(String::new())), ("trace_ref".into(), Parameter::Text(first_ref.into())), ("span".into(), Parameter::Text("root".into())), ("quote".into(), Parameter::Text(opposite.into())), @@ -1339,7 +1341,7 @@ async fn lens_selection_pages_without_losing_or_repeating_runs( #[case::traces("traces", 9)] #[case::requests("requests", 3)] #[tokio::test] -async fn lens_content_keeps_original_span_and_request_timestamps( +async fn lens_content_keeps_original_timestamps_with_start_time_slack( #[future(awt)] database: TestResult, #[case] source: &str, #[case] precision: usize, @@ -1379,6 +1381,30 @@ async fn lens_content_keeps_original_span_and_request_timestamps( ) .await?; let connection = Connection::configured(&database.url, "trace_test", "default", "")?; + let start_time_body = execute_read( + &database.client, + &connection, + "SELECT toString(fromUnixTimestamp64Nano({timestamp:Int64})) AS start_time FORMAT JSON", + &BTreeMap::from([( + "timestamp".into(), + Parameter::Integer(root_start + 86_400_000_000_000), + )]), + ) + .await?; + let start_time: serde_json::Value = serde_json::from_str(&start_time_body)?; + let start_time = start_time["data"][0]["start_time"] + .as_str() + .ok_or("start time missing")? + .to_owned(); + let parsed_time_body = execute_read( + &database.client, + &connection, + "SELECT toString(parseDateTime64BestEffortOrZero({start_time:String}, 9)) AS start_time FORMAT JSON", + &BTreeMap::from([("start_time".into(), Parameter::Text(start_time.clone()))]), + ) + .await?; + let parsed_time: serde_json::Value = serde_json::from_str(&parsed_time_body)?; + assert_eq!(parsed_time["data"][0]["start_time"], start_time); let parameters = BTreeMap::from([ ("source".into(), Parameter::Text(source.into())), ("all_teams".into(), Parameter::Integer(0)), @@ -1386,6 +1412,7 @@ async fn lens_content_keeps_original_span_and_request_timestamps( ("record_team".into(), Parameter::Text("team".into())), ("key_hash".into(), Parameter::Text(String::new())), ("trace_ref".into(), Parameter::Text(String::new())), + ("start_time".into(), Parameter::Text(start_time)), ("id".into(), Parameter::Text("run".into())), ("cursor".into(), Parameter::Text(String::new())), ("offset".into(), Parameter::Integer(1)), @@ -1453,6 +1480,7 @@ async fn lens_content_keeps_output_visible_after_long_input( ("record_team".into(), Parameter::Text("team".into())), ("key_hash".into(), Parameter::Text(String::new())), ("trace_ref".into(), Parameter::Text(String::new())), + ("start_time".into(), Parameter::Text(String::new())), ("id".into(), Parameter::Text("request".into())), ("cursor".into(), Parameter::Text(String::new())), ("offset".into(), Parameter::Integer(1)), diff --git a/litellm/integrations/otel/model/spans.py b/litellm/integrations/otel/model/spans.py index 7cdb375f9d3..762efa46ebb 100644 --- a/litellm/integrations/otel/model/spans.py +++ b/litellm/integrations/otel/model/spans.py @@ -354,6 +354,8 @@ _PRISMA_MODELS: Final[frozenset[str]] = frozenset( "LiteLLM_LensReview", "LiteLLM_LensIngestionKey", "LiteLLM_LensWorker", + "LiteLLM_LensSignalConfig", + "LiteLLM_LensTraceSignal", ) ) PRISMA_RELATIONS: Final[frozenset[str]] = _PRISMA_MODELS | _PRISMA_VIEWS diff --git a/litellm/litellm_core_utils/prompt_templates/factory.py b/litellm/litellm_core_utils/prompt_templates/factory.py index 973a2787b6a..4580f9bd01b 100644 --- a/litellm/litellm_core_utils/prompt_templates/factory.py +++ b/litellm/litellm_core_utils/prompt_templates/factory.py @@ -5188,7 +5188,7 @@ def function_call_prompt( messages: list[dict[str, object]], functions: list[object], ) -> list[dict[str, object]]: - function_prompt = """Produce JSON OUTPUT ONLY! Adhere to this format {"name": "function_name", "arguments":{"argument_name": "argument_value"}} The following functions are available to you:""" + function_prompt = """To call a function, reply with JSON ONLY in this format {"name": "function_name", "arguments":{"argument_name": "argument_value"}}. Once a function result answers the request, reply to the user in plain text instead of calling a function again. The following functions are available to you:""" for function in functions: function_prompt += f"""\n{function}\n""" diff --git a/litellm/llms/base_llm/chat/transformation.py b/litellm/llms/base_llm/chat/transformation.py index 50825720331..8275829765b 100644 --- a/litellm/llms/base_llm/chat/transformation.py +++ b/litellm/llms/base_llm/chat/transformation.py @@ -65,6 +65,26 @@ class BaseLLMException(Exception): super().__init__(self.message) # Call the base class constructor with the parameters it needs +_NO_ATTRIBUTION_HEADERS: Final[Mapping[str, str]] = types.MappingProxyType({}) + + +def with_attribution_headers( + attribution_headers: Mapping[str, str], + headers: dict[str, str] | None, # mutable-ok: returned as-is when there is nothing to add +) -> dict[str, str] | None: # mutable-ok: becomes the request's outbound headers + """ + `headers` plus any attribution header the caller didn't already set (names + compared case-insensitively). Builds a new dict; `headers` is never mutated. + """ + if not attribution_headers: + return headers + caller_names: Final = {name.lower() for name in headers or {}} + return { + **{name: value for name, value in attribution_headers.items() if name.lower() not in caller_names}, + **(headers or {}), + } + + class BaseConfig(ABC): def __init__(self): pass @@ -89,6 +109,15 @@ class BaseConfig(ABC): and not callable(v) # Filter out any callable objects including mocks } + def get_attribution_headers(self) -> Mapping[str, str]: + """ + Headers that tell the provider a request came through LiteLLM. + + Sent by default on every request; a caller header with the same name + (any casing) wins. Override in a provider config to opt in. + """ + return _NO_ATTRIBUTION_HEADERS + def get_json_schema_from_pydantic_object(self, response_format: type[BaseModel] | dict | None) -> dict | None: return type_to_response_format_param(response_format=response_format) diff --git a/litellm/llms/novita/chat/transformation.py b/litellm/llms/novita/chat/transformation.py index acdfa7e8790..f1e1e3c91b1 100644 --- a/litellm/llms/novita/chat/transformation.py +++ b/litellm/llms/novita/chat/transformation.py @@ -6,11 +6,20 @@ Calls done in OpenAI/openai.py as Novita AI is openai-compatible. Docs: https://novita.ai/docs/guides/llm-api """ +from collections.abc import Mapping +from types import MappingProxyType +from typing import Final + from ....types.llms.openai import AllMessageValues from ...openai.chat.gpt_transformation import OpenAIGPTConfig +_NOVITA_ATTRIBUTION_HEADERS: Final[Mapping[str, str]] = MappingProxyType({"X-Novita-Source": "litellm"}) + class NovitaConfig(OpenAIGPTConfig): + def get_attribution_headers(self) -> Mapping[str, str]: + return _NOVITA_ATTRIBUTION_HEADERS + def validate_environment( self, headers: dict, @@ -27,5 +36,6 @@ class NovitaConfig(OpenAIGPTConfig): ) headers["Authorization"] = f"Bearer {api_key}" headers["Content-Type"] = "application/json" - headers["X-Novita-Source"] = "litellm" + if not any(name.lower() == "x-novita-source" for name in headers): + headers["X-Novita-Source"] = "litellm" return headers diff --git a/litellm/llms/perplexity/chat/transformation.py b/litellm/llms/perplexity/chat/transformation.py index 2894218da67..5b85e1e8ae6 100644 --- a/litellm/llms/perplexity/chat/transformation.py +++ b/litellm/llms/perplexity/chat/transformation.py @@ -2,6 +2,8 @@ Translate from OpenAI's `/v1/chat/completions` to Perplexity's `/v1/chat/completions` """ +from collections.abc import Mapping +from types import MappingProxyType from typing import TYPE_CHECKING, Final import httpx @@ -17,12 +19,17 @@ from litellm.types.utils import ModelResponse, PromptTokensDetailsWrapper, Usage if TYPE_CHECKING: from litellm.litellm_core_utils.tokenizer import Encoding as Tokenizer +_PERPLEXITY_ATTRIBUTION_HEADERS: Final[Mapping[str, str]] = MappingProxyType({"X-Pplx-Integration": "litellm"}) + class PerplexityChatConfig(OpenAIGPTConfig): @property def custom_llm_provider(self) -> str | None: return "perplexity" + def get_attribution_headers(self) -> Mapping[str, str]: + return _PERPLEXITY_ATTRIBUTION_HEADERS + def _get_openai_compatible_provider_info( self, api_base: str | None, api_key: str | None ) -> tuple[str | None, str | None]: diff --git a/litellm/main.py b/litellm/main.py index 519b8aad0ef..37b290f51ba 100644 --- a/litellm/main.py +++ b/litellm/main.py @@ -118,6 +118,7 @@ from litellm.llms.base_llm import BaseConfig, BaseImageGenerationConfig from litellm.llms.base_llm.base_model_iterator import ( convert_model_response_to_streaming, ) +from litellm.llms.base_llm.chat.transformation import with_attribution_headers from litellm.llms.bedrock.common_utils import ( BedrockModelInfo, bedrock_route_for_request, @@ -2634,6 +2635,11 @@ def _complete_custom_openai( ) headers = headers or litellm.headers + outbound_headers: Final = ( + headers + if provider_config is None + else with_attribution_headers(provider_config.get_attribution_headers(), headers) + ) # Add GitHub Copilot headers (same as /responses endpoint does) if custom_llm_provider == "github_copilot": @@ -2685,7 +2691,7 @@ def _complete_custom_openai( acompletion=acompletion, stream=stream, api_key=api_key, - headers=headers, + headers=outbound_headers, client=client, provider_config=provider_config, ) @@ -2693,7 +2699,7 @@ def _complete_custom_openai( response = openai_chat_completions.completion( model=model, messages=messages, - headers=headers, + headers=outbound_headers, model_response=model_response, print_verbose=print_verbose, api_key=api_key, @@ -2716,7 +2722,7 @@ def _complete_custom_openai( input=messages, api_key=api_key, original_response=str(e), - additional_args={"headers": headers}, + additional_args={"headers": outbound_headers}, ) raise e @@ -2726,7 +2732,7 @@ def _complete_custom_openai( input=messages, api_key=api_key, original_response=response, - additional_args={"headers": headers}, + additional_args={"headers": outbound_headers}, ) return response # pyright: ignore[reportReturnType] # provider SDK return type is broader than the dispatch contract diff --git a/litellm/proxy/lens/endpoints.py b/litellm/proxy/lens/endpoints.py index 9e655e83645..1b4d56e5900 100644 --- a/litellm/proxy/lens/endpoints.py +++ b/litellm/proxy/lens/endpoints.py @@ -62,6 +62,8 @@ from litellm.proxy.lens.models import ( from litellm.proxy.lens.release import PROTOCOL_VERSION, release_tag, worker_image from litellm.proxy.lens.repository import DueLens, LensRepository, WriterDatabase from litellm.proxy.lens.reviews import criteria_key +from litellm.proxy.lens.signal_repository import SignalRepository +from litellm.proxy.lens.signals import SignalConfig, TraceSignals, trace_signals from litellm.proxy.lens.sources import ActivityAvailability, SourceReader, Storage, parse_execution from litellm.proxy.lens.state import ( can_access, @@ -79,6 +81,7 @@ from litellm.proxy.lens.state import ( summarized, ) from litellm.proxy.tracing_runtime import provide_storage +from litellm.router import Router from litellm.tracing.remote import LensConnection, bounded_response from litellm.types.llms.base import LiteLLMBaseModel @@ -112,6 +115,14 @@ def repository() -> LensRepository: return LensRepository(WriterDatabase(writer_wrapper(prisma_client.db))) +def signals_repository() -> SignalRepository: + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + raise HTTPException(503, "Lens needs a connected Postgres database") + return SignalRepository(WriterDatabase(writer_wrapper(prisma_client.db))) + + def source_reader(storage: Storage | None) -> SourceReader: if storage is None: raise HTTPException( @@ -129,6 +140,20 @@ def user_scope(auth: UserAPIKeyAuth, write: bool = False) -> Scope: raise HTTPException(403, "Lens requires proxy administrator access") +def validate_signal_model(config: SignalConfig, llm_router: Router | None) -> None: + if not config.model: + return + message: Final = "Choose a System 1 model (evaluation mode) configured on this proxy" + if llm_router is None: + raise HTTPException(400, message) + try: + model_group: Final = llm_router.get_model_group_info(model_group=config.model) + except Exception as error: + raise HTTPException(400, message) from error + if model_group is None or model_group.mode != "evaluation": + raise HTTPException(400, message) + + async def get_lens(lens_id: str, scope: Scope) -> Lens: lens: Final = await repository().get(lens_id) if lens is None or not can_access(scope, lens.scope): @@ -367,6 +392,39 @@ async def list_agents(auth: Auth, storage: StorageDep) -> tuple[str, ...]: return await source_reader(storage).agents(scope) if storage is not None else () +@router.get("/signals", response_model=SignalConfig) +async def get_signals(auth: Auth) -> SignalConfig: + user_scope(auth) + return await signals_repository().get_config() + + +@router.put("/signals", response_model=SignalConfig) +async def put_signals(body: SignalConfig, auth: Auth) -> SignalConfig: + user_scope(auth, write=True) + from litellm.proxy.proxy_server import llm_router + + validate_signal_model(body, llm_router) + await signals_repository().save_config(body) + return body + + +@router.post("/traces/signals", response_model=tuple[TraceSignals, ...]) +async def trace_signal_statuses(body: TraceFindingsRequest, auth: Auth) -> tuple[TraceSignals, ...]: + user_scope(auth) + repo: Final = signals_repository() + config: Final = await repo.get_config() + existing: Final = await repo.traces(body.traces) + rows: Final = MappingProxyType({(row.trace_id, row.trace_ref): row for row in existing}) + return tuple( + trace_signals( + trace, + rows.get((trace.trace_id, trace.trace_ref)), + config, + ) + for trace in body.traces + ) + + @router.post("/traces/findings", response_model=tuple[TraceFindingCount, ...]) async def trace_findings(body: TraceFindingsRequest, auth: Auth) -> tuple[TraceFindingCount, ...]: user_scope(auth) diff --git a/litellm/proxy/lens/signal_repository.py b/litellm/proxy/lens/signal_repository.py new file mode 100644 index 00000000000..462cbdf13f4 --- /dev/null +++ b/litellm/proxy/lens/signal_repository.py @@ -0,0 +1,142 @@ +import json +from datetime import datetime +from typing import Final + +from pydantic import TypeAdapter + +from litellm.proxy.lens.models import Execution, TraceIdentity +from litellm.proxy.lens.repository import Database, Row +from litellm.proxy.lens.signals import ( + SIGNAL_RECLASSIFY_AFTER, + SIGNAL_RETRY_FAILED_AFTER, + SignalAttempt, + SignalConfig, + StoredTraceSignal, +) + +_ROWS: Final[TypeAdapter[tuple[Row, ...]]] = TypeAdapter(tuple[Row, ...]) + + +class SignalRepository: + def __init__(self, db: Database) -> None: + self.db: Final = db + + async def get_config(self) -> SignalConfig: + rows: Final = _ROWS.validate_python( + await self.db.query_raw('SELECT data FROM "LiteLLM_LensSignalConfig" WHERE id=$1', "global") + ) + return SignalConfig() if not rows else SignalConfig.model_validate(rows[0].data) + + async def save_config(self, config: SignalConfig) -> None: + await self.db.execute_raw( + """INSERT INTO "LiteLLM_LensSignalConfig" (id, data) + VALUES ($1, $2::jsonb) + ON CONFLICT (id) DO UPDATE SET data=EXCLUDED.data""", + "global", + json.dumps(config.model_dump(mode="json")), + ) + + async def traces(self, identities: tuple[TraceIdentity, ...]) -> tuple[StoredTraceSignal, ...]: + if not identities: + return () + payload: Final = json.dumps( + tuple({"trace_id": trace.trace_id, "trace_ref": trace.trace_ref} for trace in identities) + ) + rows: Final = _ROWS.validate_python( + await self.db.query_raw( + """SELECT jsonb_build_object( + 'trace_id', trace_id, + 'trace_ref', trace_ref, + 'config_key', config_key, + 'span_count', span_count, + 'claimed_until', claimed_until, + 'classified_at', classified_at, + 'data', data + ) AS data + FROM "LiteLLM_LensTraceSignal" + WHERE (trace_id, trace_ref) IN ( + SELECT trace_id, trace_ref FROM jsonb_to_recordset($1::jsonb) AS requested( + trace_id text, trace_ref text + ) + )""", + payload, + ) + ) + return tuple(StoredTraceSignal.model_validate(row.data) for row in rows) + + async def claim( + self, + execution: Execution, + config: SignalConfig, + claimed_until: datetime, + now: datetime, + ) -> bool: + data: Final = json.dumps({"status": "pending", "scores": {}, "model": config.model, "error": ""}) + rows: Final = _ROWS.validate_python( + await self.db.query_raw( + """INSERT INTO "LiteLLM_LensTraceSignal" AS stored + (trace_id, trace_ref, config_key, span_count, claimed_until, classified_at, data) + VALUES ($1, $2, $3, $4, $5::timestamp, NULL, $6::jsonb) + ON CONFLICT (trace_id, trace_ref) DO UPDATE SET + config_key=EXCLUDED.config_key, + span_count=EXCLUDED.span_count, + claimed_until=EXCLUDED.claimed_until, + classified_at=NULL, + data=EXCLUDED.data + WHERE (stored.claimed_until IS NULL OR stored.claimed_until < $7::timestamp) + AND ( + stored.config_key IS DISTINCT FROM EXCLUDED.config_key + OR ( + stored.data->>'status'='pending' + AND stored.claimed_until < $7::timestamp + ) + OR ( + EXCLUDED.span_count > stored.span_count + AND stored.classified_at < $8::timestamp + ) + OR ( + stored.data->>'status'='failed' + AND stored.classified_at < $9::timestamp + ) + ) + RETURNING jsonb_build_object('trace_id', trace_id) AS data""", + execution.trace_id, + execution.trace_ref, + config.key(), + execution.span_count, + claimed_until, + data, + now, + now - SIGNAL_RECLASSIFY_AFTER, + now - SIGNAL_RETRY_FAILED_AFTER, + ) + ) + return bool(rows) + + async def store( + self, + execution: Execution, + config: SignalConfig, + claimed_until: datetime, + classified_at: datetime, + attempt: SignalAttempt, + ) -> None: + payload: Final = json.dumps( + { + "status": attempt.status, + "scores": dict(attempt.scores), + "model": attempt.model, + "error": attempt.error, + } + ) + await self.db.execute_raw( + """UPDATE "LiteLLM_LensTraceSignal" + SET classified_at=$1::timestamp, claimed_until=NULL, data=$2::jsonb + WHERE trace_id=$3 AND trace_ref=$4 AND config_key=$5 AND claimed_until=$6::timestamp""", + classified_at, + payload, + execution.trace_id, + execution.trace_ref, + config.key(), + claimed_until, + ) diff --git a/litellm/proxy/lens/signals.py b/litellm/proxy/lens/signals.py new file mode 100644 index 00000000000..cfe273eb6af --- /dev/null +++ b/litellm/proxy/lens/signals.py @@ -0,0 +1,586 @@ +import asyncio +import hashlib +import json +from collections.abc import Callable, Mapping +from datetime import datetime, timedelta, timezone +from itertools import accumulate +from types import MappingProxyType +from typing import Annotated, Final, Literal, Protocol, TypeAlias + +from pydantic import ConfigDict, Field, JsonValue, ValidationError, field_validator, model_validator + +from litellm.integrations.clickhouse.context import lens_analysis +from litellm.litellm_core_utils.initialize_dynamic_callback_params import inherit_message_logging_privacy +from litellm.litellm_core_utils.secret_redaction import redact_internal_details +from litellm.proxy.lens.models import ActivitySelection, Execution, Record, Scope, TraceIdentity +from litellm.proxy.lens.sources import SourceReader, Storage + +SIGNAL_INTERVAL_SECONDS: Final = 60 +SIGNAL_PAGE_SIZE: Final = 100 +SIGNAL_MAX_PER_TICK: Final = 50 +SIGNAL_CONCURRENCY: Final = 8 +SIGNAL_CLAIM_LEASE: Final = timedelta(minutes=5) +SIGNAL_RECLASSIFY_AFTER: Final = timedelta(minutes=5) +SIGNAL_RETRY_FAILED_AFTER: Final = timedelta(minutes=30) +SIGNAL_MAX_CONTENT_PAGES: Final = 3 +SIGNAL_PART_MAX_CHARS: Final = 2000 +SIGNAL_PART_HEAD_CHARS: Final = 800 +SIGNAL_PART_TAIL_CHARS: Final = 1200 +SIGNAL_TRANSCRIPT_MAX_CHARS: Final = 40000 +SIGNAL_TRANSCRIPT_HEAD_CHARS: Final = 15000 +SIGNAL_TRANSCRIPT_TAIL_CHARS: Final = 25000 +SIGNAL_MAX_SCAN_PAGES: Final = 10 +SIGNAL_TASK: Final = ( + "An AI agent run recorded as a trace. Judge only what the user and the agent said and did in these steps." +) + + +class Signal(Record): + id: str = Field(pattern=r"^[a-z][a-z0-9_]{0,63}$") + name: str = Field(min_length=1, max_length=60) + question: str = Field(min_length=3, max_length=500) + + +DEFAULT_SIGNALS: Final[tuple[Signal, ...]] = ( + Signal( + id="user_frustration", + name="User frustration", + question=( + "Does the user show frustration, annoyance or dissatisfaction with the agent in this run, for example " + "complaints, irritated corrections, all caps, profanity, or giving up on the task?" + ), + ), + Signal( + id="missing_capability", + name="Missing capability", + question=( + "Does the user ask for something the agent cannot do in this run, so that the agent refuses, says it " + "lacks a tool, permission, integration or data source, or fails because the capability does not exist?" + ), + ), + Signal( + id="repeated_request", + name="Repeated request", + question=( + "Does the user ask for the same thing more than once in this run, usually because the agent did not " + "deliver it the first time?" + ), + ), +) + + +class SignalConfig(Record): + model: str = "" + threshold: float = Field(default=0.5, ge=0.05, le=0.95, allow_inf_nan=False) + signals: tuple[Signal, ...] = DEFAULT_SIGNALS + + @model_validator(mode="after") + def validate_signals(self) -> "SignalConfig": + if len(self.signals) > 20: + raise ValueError("A maximum of 20 signals is allowed") + if len(frozenset(signal.id for signal in self.signals)) != len(self.signals): + raise ValueError("Signal IDs must be unique") + return self + + @property + def enabled(self) -> bool: + return bool(self.model) and bool(self.signals) + + def key(self) -> str: + payload: Final = json.dumps( + { + "model": self.model, + "signals": tuple({"id": signal.id, "question": signal.question} for signal in self.signals), + }, + ensure_ascii=False, + sort_keys=True, + separators=(",", ":"), + ) + return hashlib.sha256(payload.encode()).hexdigest() + + +Score: TypeAlias = Annotated[float, Field(ge=0, le=1, allow_inf_nan=False)] + + +class SignalFlag(Record): + signal_id: str + name: str + score: Score + + +class TraceSignals(TraceIdentity): + status: Literal["unclassified", "pending", "classified", "failed"] + flags: tuple[SignalFlag, ...] = () + model: str = "" + classified_at: datetime | None = None + + +class SignalStep(Record): + kind: str + name: str + content: str + + +class SignalData(Record): + model_config = ConfigDict(extra="ignore") + + status: Literal["pending", "classified", "failed"] = "pending" + scores: Mapping[str, Score] = Field(default_factory=lambda: MappingProxyType({})) + model: str = "" + error: str = "" + + +class SignalAttempt(Record): + status: Literal["classified", "failed"] + scores: Mapping[str, Score] = Field(default_factory=lambda: MappingProxyType({})) + model: str + error: str = "" + + +class StoredTraceSignal(Record): + trace_id: str + trace_ref: str = "" + config_key: str + span_count: int + claimed_until: datetime | None = None + classified_at: datetime | None = None + data: JsonValue + + @field_validator("claimed_until", "classified_at") + @classmethod + def normalize_database_timestamp(cls, value: datetime | None) -> datetime | None: + if value is not None and value.tzinfo is None: + return value.replace(tzinfo=timezone.utc) + return value + + +class NoulAnswer(Record): + model_config = ConfigDict(extra="ignore", allow_inf_nan=False, from_attributes=True) + + type: Literal["noul"] + noul: float = Field(ge=0, le=1, allow_inf_nan=False) + + +class DecisionsOutput(Record): + model_config = ConfigDict(extra="ignore", from_attributes=True) + + answers: Mapping[str, object] + + +DecisionState: TypeAlias = Mapping[str, object] +DecisionQuestions: TypeAlias = Mapping[str, Mapping[str, str]] +Clock: TypeAlias = Callable[[], datetime] +RouterReady: TypeAlias = Callable[[], bool] + + +class DecisionsCall(Protocol): + async def __call__( + self, + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: ... + + +class SignalRepositoryProtocol(Protocol): + async def get_config(self) -> SignalConfig: ... + async def traces(self, identities: tuple[TraceIdentity, ...]) -> tuple[StoredTraceSignal, ...]: ... + async def claim( + self, + execution: Execution, + config: SignalConfig, + claimed_until: datetime, + now: datetime, + ) -> bool: ... + async def store( + self, + execution: Execution, + config: SignalConfig, + claimed_until: datetime, + classified_at: datetime, + attempt: SignalAttempt, + ) -> None: ... + + +def signal_identity(trace: TraceIdentity | StoredTraceSignal | Execution) -> tuple[str, str]: + return trace.trace_id, trace.trace_ref + + +def candidate( + trace: Execution, + existing: StoredTraceSignal | None, + config_key: str, + now: datetime, +) -> bool: + if existing is None: + return True + if existing.claimed_until is not None and existing.claimed_until > now: + return False + if existing.config_key != config_key: + return True + status: Final = existing.data.get("status") if isinstance(existing.data, dict) else "" + if status == "pending": + return existing.claimed_until is not None and existing.claimed_until <= now + if existing.span_count > trace.span_count: + return False + if existing.span_count < trace.span_count: + return existing.classified_at is not None and existing.classified_at < now - SIGNAL_RECLASSIFY_AFTER + return ( + status == "failed" + and existing.classified_at is not None + and existing.classified_at < now - SIGNAL_RETRY_FAILED_AFTER + ) + + +def _take_head(steps: tuple[SignalStep, ...], remaining: int) -> tuple[SignalStep, ...]: + if remaining <= 0: + return () + cumulative_lengths: Final = tuple(accumulate(len(step.content) for step in steps)) + boundary: Final = next((index for index, total in enumerate(cumulative_lengths) if total >= remaining), None) + if boundary is None: + return steps + preceding: Final = steps[:boundary] + last: Final = steps[boundary] + used: Final = cumulative_lengths[boundary - 1] if boundary > 0 else 0 + last_length: Final = remaining - used + return ( + *preceding, + last + if last_length == len(last.content) + else last.model_copy(update=MappingProxyType({"content": last.content[:last_length]})), + ) + + +def _take_tail(steps: tuple[SignalStep, ...], remaining: int) -> tuple[SignalStep, ...]: + if remaining <= 0: + return () + reversed_steps: Final = tuple(reversed(steps)) + cumulative_lengths: Final = tuple(accumulate(len(step.content) for step in reversed_steps)) + boundary: Final = next((index for index, total in enumerate(cumulative_lengths) if total >= remaining), None) + if boundary is None: + return steps + preceding: Final = reversed_steps[:boundary] + last: Final = reversed_steps[boundary] + used: Final = cumulative_lengths[boundary - 1] if boundary > 0 else 0 + last_length: Final = remaining - used + selected: Final = ( + *preceding, + last + if last_length == len(last.content) + else last.model_copy(update=MappingProxyType({"content": last.content[-last_length:]})), + ) + return tuple(reversed(selected)) + + +def _bounded_steps(steps: tuple[SignalStep, ...]) -> tuple[SignalStep, ...]: + if sum(len(step.content) for step in steps) <= SIGNAL_TRANSCRIPT_MAX_CHARS: + return steps + head: Final = _take_head(steps, SIGNAL_TRANSCRIPT_HEAD_CHARS) + tail: Final = _take_tail(steps, SIGNAL_TRANSCRIPT_TAIL_CHARS) + omitted_count: Final = len(steps) - len(head) - len(tail) + marker: Final = SignalStep(kind="omitted", name="", content=f"{omitted_count} steps omitted") + return (*head, marker, *tail) + + +def _part_excerpt(content: str) -> str: + if len(content) <= SIGNAL_PART_MAX_CHARS: + return content + omitted: Final = len(content) - SIGNAL_PART_MAX_CHARS + marker: Final = f"\n[... {omitted} characters omitted ...]\n" + return f"{content[:SIGNAL_PART_HEAD_CHARS]}{marker}{content[-SIGNAL_PART_TAIL_CHARS:]}" + + +async def _content_pages( + reader: SourceReader, + scope: Scope, + execution: Execution, + cursor: str, + pages_left: int, +) -> tuple[SignalStep, ...]: + if pages_left == 0: + return () + content: Final = await reader.content(scope, execution, cursor) + current: Final = tuple( + SignalStep(kind=part.kind, name=part.name, content=_part_excerpt(part.content)) for part in content.parts + ) + rest: Final = ( + await _content_pages(reader, scope, execution, content.next_cursor, pages_left - 1) + if content.next_cursor is not None + else () + ) + return (*current, *rest) + + +async def signal_state(reader: SourceReader, scope: Scope, execution: Execution) -> DecisionState: + steps: Final = _bounded_steps(await _content_pages(reader, scope, execution, "", SIGNAL_MAX_CONTENT_PAGES)) + return { + "task": SIGNAL_TASK, + "steps": tuple(step.model_dump(mode="json") for step in steps), + } + + +def _noul_score(value: object) -> float | None: + try: + return NoulAnswer.model_validate(value).noul + except ValidationError: + return None + + +class SignalClassifier: + def __init__(self, reader: SourceReader, completion: DecisionsCall, clock: Clock) -> None: + self.reader: Final = reader + self.completion: Final = completion + self.clock: Final = clock + + async def classify(self, scope: Scope, execution: Execution, config: SignalConfig) -> SignalAttempt: + try: + state: Final = await signal_state(self.reader, scope, execution) + questions: Final = { + signal.id: {"type": "noul", "instructions": signal.question} for signal in config.signals + } + with lens_analysis(), inherit_message_logging_privacy(True): + response: Final = await self.completion( + model=config.model, + state=state, + questions=questions, + timeout=60, + metadata={"tags": ["litellm-lens-signals"]}, + ) + output: Final = DecisionsOutput.model_validate(response) + scores: Final = MappingProxyType( + { + signal.id: score + for signal in config.signals + if (score := _noul_score(output.answers.get(signal.id))) is not None + } + ) + if len(scores) != len(config.signals): + return SignalAttempt( + status="failed", + scores=scores, + model=config.model, + error="Decisions response omitted a configured noul answer", + ) + return SignalAttempt(status="classified", scores=scores, model=config.model) + except Exception as error: + detail: Final = redact_internal_details(str(error))[:300] + return SignalAttempt(status="failed", model=config.model, error=detail) + + +def trace_signals( + trace: TraceIdentity, + existing: StoredTraceSignal | None, + config: SignalConfig, +) -> TraceSignals: + if existing is None or existing.config_key != config.key(): + return TraceSignals(trace_id=trace.trace_id, trace_ref=trace.trace_ref, status="unclassified") + data: Final = SignalData.model_validate(existing.data) + if data.status == "pending": + return TraceSignals( + trace_id=trace.trace_id, + trace_ref=trace.trace_ref, + status="pending", + model=data.model, + ) + if data.status == "failed" or data.error: + return TraceSignals( + trace_id=trace.trace_id, + trace_ref=trace.trace_ref, + status="failed", + model=data.model, + classified_at=existing.classified_at, + ) + flags: Final = tuple( + sorted( + ( + SignalFlag(signal_id=signal.id, name=signal.name, score=data.scores[signal.id]) + for signal in config.signals + if signal.id in data.scores and data.scores[signal.id] >= config.threshold + ), + key=lambda flag: flag.score, + reverse=True, + ) + ) + return TraceSignals( + trace_id=trace.trace_id, + trace_ref=trace.trace_ref, + status="classified", + flags=flags, + model=data.model, + classified_at=existing.classified_at, + ) + + +async def _process_claimed( + classifier: SignalClassifier, + repository: SignalRepositoryProtocol, + scope: Scope, + execution: Execution, + config: SignalConfig, + claimed_until: datetime, +) -> None: + from litellm._logging import verbose_proxy_logger + + attempt: Final = await classifier.classify(scope, execution, config) + try: + await repository.store(execution, config, claimed_until, classifier.clock(), attempt) + except Exception as error: + verbose_proxy_logger.error("Lens signal result could not be stored: %s", redact_internal_details(str(error))) + + +class _SignalScan: + def __init__( + self, + reader: SourceReader, + repository: SignalRepositoryProtocol, + scope: Scope, + config: SignalConfig, + now: datetime, + cursor: str, + limit: int, + ) -> None: + self.reader: Final = reader + self.repository: Final = repository + self.scope: Final = scope + self.config: Final = config + self.now: Final = now + self.cursor: str = cursor + self.limit: Final = limit + self.executions: tuple[Execution, ...] = () + self.finished: bool = False + + async def _read_page(self, start: int, end: int) -> tuple[tuple[Execution, ...], str | None]: + page_cursor: Final = self.cursor + sample: Final = await self.reader.sample( + self.scope, + ActivitySelection(source="traces"), + start, + end, + page_size=SIGNAL_PAGE_SIZE, + cursor=page_cursor, + ) + identities: Final = tuple( + TraceIdentity(trace_id=trace.trace_id, trace_ref=trace.trace_ref) for trace in sample.executions + ) + existing_rows: Final = await self.repository.traces(identities) + existing: Final = MappingProxyType({signal_identity(row): row for row in existing_rows}) + remaining: Final = self.limit - len(self.executions) + all_eligible: Final = tuple( + execution + for execution in sample.executions + if candidate(execution, existing.get(signal_identity(execution)), self.config.key(), self.now) + ) + eligible: Final = all_eligible[:remaining] + next_cursor: Final = page_cursor if len(all_eligible) > remaining else sample.next_cursor + return eligible, next_cursor + + async def run(self) -> tuple[tuple[Execution, ...], str]: + start: Final = int((self.now - timedelta(hours=24)).timestamp() * 1000) + end: Final = int((self.now - timedelta(minutes=2)).timestamp() * 1000) + for _ in range(SIGNAL_MAX_SCAN_PAGES): + if self.finished or len(self.executions) >= self.limit: + break + eligible, next_cursor = await self._read_page(start, end) + self.executions = (*self.executions, *eligible) + if next_cursor is None: + self.cursor = "" + self.finished = True + else: + self.cursor = next_cursor + return self.executions, self.cursor + + +async def _scan_pages( + reader: SourceReader, + repository: SignalRepositoryProtocol, + scope: Scope, + config: SignalConfig, + now: datetime, + cursor: str, + remaining: int, +) -> tuple[tuple[Execution, ...], str]: + if remaining <= 0: + return (), cursor + scan: Final = _SignalScan(reader, repository, scope, config, now, cursor, remaining) + return await scan.run() + + +async def run_signal_tick( + storage: Storage, + repository: SignalRepositoryProtocol | None, + completion: DecisionsCall | None, + clock: Clock, + router_ready: RouterReady = lambda: True, + cursor: str = "", +) -> str: + if repository is None or completion is None or not router_ready(): + return cursor + now: Final = clock() + config: Final = await repository.get_config() + if not config.enabled: + return cursor + reader: Final = SourceReader(storage) + scope: Final = Scope(all_teams=True) + candidates: Final = await _scan_pages( + reader, + repository, + scope, + config, + now, + cursor, + SIGNAL_MAX_PER_TICK, + ) + executions, next_cursor = candidates + classifier: Final = SignalClassifier(reader, completion, clock) + semaphore: Final = asyncio.Semaphore(SIGNAL_CONCURRENCY) + + async def process(execution: Execution) -> None: + from litellm._logging import verbose_proxy_logger + + async with semaphore: + claimed_at: Final = classifier.clock() + claimed_until: Final = claimed_at + SIGNAL_CLAIM_LEASE + try: + claimed: Final = await repository.claim(execution, config, claimed_until, claimed_at) + except Exception as error: + verbose_proxy_logger.error("Lens signal claim failed: %s", redact_internal_details(str(error))) + return + if not claimed: + return + await _process_claimed(classifier, repository, scope, execution, config, claimed_until) + + await asyncio.gather(*(process(execution) for execution in executions)) + return next_cursor + + +class _SignalLoopState: + def __init__(self) -> None: + self.cursor: str = "" + + +async def run_signal_loop( + storage: Storage, + repository: SignalRepositoryProtocol | None, + completion: DecisionsCall | None, + clock: Clock = lambda: datetime.now(timezone.utc), + router_ready: RouterReady = lambda: True, +) -> None: + from litellm._logging import verbose_proxy_logger + + state: Final = _SignalLoopState() + while True: + try: + state.cursor = await run_signal_tick( + storage, + repository, + completion, + clock, + router_ready, + cursor=state.cursor, + ) + except Exception as error: + verbose_proxy_logger.error("Lens signal tick failed: %s", redact_internal_details(str(error))) + await asyncio.sleep(SIGNAL_INTERVAL_SECONDS) diff --git a/litellm/proxy/lens/sources.py b/litellm/proxy/lens/sources.py index 015be69ecca..b1d142a6748 100644 --- a/litellm/proxy/lens/sources.py +++ b/litellm/proxy/lens/sources.py @@ -142,6 +142,7 @@ class SourceReader: id=execution.trace_id, trace_ref=execution.trace_ref, record_team=execution.team_id, + start_time=execution.start_time, cursor=cursor, offset=offset + 1, ) @@ -175,6 +176,7 @@ class SourceReader: id=execution.trace_id, trace_ref=execution.trace_ref, record_team=execution.team_id, + start_time=execution.start_time, span=evidence.span_id, quote=evidence.quote, ) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index eeb099866d5..686fe71bbd0 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -581,6 +581,14 @@ from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger, run_sp from litellm.proxy.image_endpoints.endpoints import router as image_router from litellm.proxy.lens.dataset_endpoints import router as lens_dataset_router from litellm.proxy.lens.endpoints import router as lens_router +from litellm.proxy.lens.repository import WriterDatabase +from litellm.proxy.lens.signal_repository import SignalRepository +from litellm.proxy.lens.signals import ( + DecisionQuestions, + DecisionsCall, + DecisionState, + run_signal_loop, +) from litellm.proxy.list_api.common import ( ManagementProblem, problem_response, @@ -1275,6 +1283,27 @@ async def _connect_to_count_stored_values() -> SupportsRawQueries: return client.writer_db +async def _call_current_lens_signal_router( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], +) -> object: + current_router: Final = llm_router + if current_router is None: + raise RuntimeError("The proxy router is not initialized") + decisions: Final[DecisionsCall] = cast(DecisionsCall, current_router.adecisions) + return await decisions( + model=model, + state=state, + questions=questions, + timeout=timeout, + metadata=metadata, + ) + + @asynccontextmanager async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState, None]: global \ @@ -1645,12 +1674,30 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[ProxyLifespanState state: Final[ProxyLifespanState] = {"tracing_receiver": receiver} from litellm.proxy.admin_mcp import admin_mcp_lifespan + signal_completion: Final[DecisionsCall] = _call_current_lens_signal_router + signal_task: Final = ( + asyncio.create_task( + run_signal_loop( + receiver.storage, + SignalRepository(WriterDatabase(writer_wrapper(prisma_client.db))), + signal_completion, + router_ready=lambda: llm_router is not None, + ) + ) + if receiver is not None and prisma_client is not None + else None + ) + try: async with AsyncExitStack() as admin_mcp_stack: try: await admin_mcp_stack.enter_async_context(admin_mcp_lifespan(app)) yield state finally: + if signal_task is not None: + signal_task.cancel() + await asyncio.gather(signal_task, return_exceptions=True) + if model_info_scheduler is not None and model_info_scheduler.running: model_info_scheduler.remove_job("refresh_model_info") if model_info_scheduler is not scheduler: @@ -13951,7 +13998,7 @@ async def run_thread( # ) # async def get_available_routes(user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth)): from litellm.llms.base_llm.base_utils import BaseTokenCounter -from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient +from litellm.proxy.db.routing_prisma_wrapper import WriterPinnedClient, writer_wrapper from litellm.repositories.config_repository import ConfigRepository from litellm.repositories.model_repository import ModelRepository from litellm.repositories.table_repositories import ( diff --git a/litellm/proxy/schema.prisma b/litellm/proxy/schema.prisma index 2f1d4aab78e..1b5d870ce1b 100644 --- a/litellm/proxy/schema.prisma +++ b/litellm/proxy/schema.prisma @@ -1982,3 +1982,20 @@ model LiteLLM_LensDataset { @@id([id, revision]) } + +model LiteLLM_LensSignalConfig { + id String @id + data Json +} + +model LiteLLM_LensTraceSignal { + trace_id String + trace_ref String @default("") + config_key String + span_count Int + claimed_until DateTime? + classified_at DateTime? + data Json + + @@id([trace_id, trace_ref]) +} diff --git a/litellm/rust_bridge/trace/generated/models.py b/litellm/rust_bridge/trace/generated/models.py index ea2c8bda648..5d84003aba2 100644 --- a/litellm/rust_bridge/trace/generated/models.py +++ b/litellm/rust_bridge/trace/generated/models.py @@ -215,6 +215,7 @@ class LensContentParams(LiteLLMBaseModel): source: ContentSource id: str record_team: str + start_time: str trace_ref: str cursor: str offset: int = Field(..., ge=0, le=4294967295) @@ -232,6 +233,7 @@ class LensEvidenceParams(LiteLLMBaseModel): source: ContentSource id: str record_team: str + start_time: str trace_ref: str span: str quote: str diff --git a/schema.prisma b/schema.prisma index 2f1d4aab78e..1b5d870ce1b 100644 --- a/schema.prisma +++ b/schema.prisma @@ -1982,3 +1982,20 @@ model LiteLLM_LensDataset { @@id([id, revision]) } + +model LiteLLM_LensSignalConfig { + id String @id + data Json +} + +model LiteLLM_LensTraceSignal { + trace_id String + trace_ref String @default("") + config_key String + span_count Int + claimed_until DateTime? + classified_at DateTime? + data Json + + @@id([trace_id, trace_ref]) +} diff --git a/scripts/trace_codegen/schemas/traces-clickhouse/LensContentParams.json b/scripts/trace_codegen/schemas/traces-clickhouse/LensContentParams.json index 5ee5ab558ce..6026ccd26e1 100644 --- a/scripts/trace_codegen/schemas/traces-clickhouse/LensContentParams.json +++ b/scripts/trace_codegen/schemas/traces-clickhouse/LensContentParams.json @@ -39,6 +39,9 @@ "source": { "$ref": "#/$defs/ContentSource" }, + "start_time": { + "type": "string" + }, "team": { "type": "string" }, @@ -53,6 +56,7 @@ "source", "id", "record_team", + "start_time", "trace_ref", "cursor", "offset" diff --git a/scripts/trace_codegen/schemas/traces-clickhouse/LensEvidenceParams.json b/scripts/trace_codegen/schemas/traces-clickhouse/LensEvidenceParams.json index 07b9c216083..dbe9b32fdd6 100644 --- a/scripts/trace_codegen/schemas/traces-clickhouse/LensEvidenceParams.json +++ b/scripts/trace_codegen/schemas/traces-clickhouse/LensEvidenceParams.json @@ -36,6 +36,9 @@ "span": { "type": "string" }, + "start_time": { + "type": "string" + }, "team": { "type": "string" }, @@ -50,6 +53,7 @@ "source", "id", "record_team", + "start_time", "trace_ref", "span", "quote" diff --git a/tests/e2e/AGENTS.md b/tests/e2e/AGENTS.md index 2e2169f8604..98a56bf1860 100644 --- a/tests/e2e/AGENTS.md +++ b/tests/e2e/AGENTS.md @@ -215,7 +215,7 @@ llm..... | rerank | images_generations | audio_speech | audio_transcriptions | moderations | realtime route : openai | azure_openai | anthropic | bedrock_converse | bedrock_invoke | vertex - | azure_foundry | cohere | together_ai + | azure_foundry | cohere | together_ai | ollama | ollama_chat (vocab varies per endpoint; messages is anthropic-format only) capability : basic | tool_use | prompt_cache_5m | vision | thinking | structured_output | service_tier | mid_conversation_system diff --git a/tests/e2e/coverage_registry/llm_conversational.yaml b/tests/e2e/coverage_registry/llm_conversational.yaml index 919884b66f0..cd41a4548ee 100644 --- a/tests/e2e/coverage_registry/llm_conversational.yaml +++ b/tests/e2e/coverage_registry/llm_conversational.yaml @@ -101,6 +101,32 @@ - {id: llm.messages.together_ai.basic.stream.works, module: llm, tier: P1, subject_endpoint: messages, route: together_ai, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together over /v1/messages streaming"} - {id: llm.messages.together_ai.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: together_ai, capability: tool_use, streaming: nonstream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together tool calls over /v1/messages"} - {id: llm.messages.together_ai.multi_turn.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: together_ai, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_together_ai_e2e.py", rationale: "Together tool result round trip over /v1/messages"} +- {id: llm.chat_completions.ollama_chat.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama_chat, capability: basic, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /chat/completions returns text and usage"} +- {id: llm.chat_completions.ollama_chat.basic.stream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama_chat, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /chat/completions streams text deltas, usage and a terminal event"} +- {id: llm.chat_completions.ollama_chat.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama_chat, capability: tool_use, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /chat/completions returns one addressable get_weather call"} +- {id: llm.chat_completions.ollama_chat.multi_turn.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama_chat, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /chat/completions tool result round trip reaches the model"} +- {id: llm.chat_completions.ollama_chat.tool_use.stream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama_chat, capability: tool_use, streaming: stream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) streams the call as tool_call deltas ending in finish_reason tool_calls, the shape coding agents like OpenCode execute"} +- {id: llm.messages.ollama_chat.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: ollama_chat, capability: basic, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /v1/messages returns text and usage"} +- {id: llm.messages.ollama_chat.basic.stream.works, module: llm, tier: P1, subject_endpoint: messages, route: ollama_chat, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /v1/messages streams text deltas, usage and a terminal event"} +- {id: llm.messages.ollama_chat.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: ollama_chat, capability: tool_use, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /v1/messages returns one addressable get_weather call"} +- {id: llm.messages.ollama_chat.multi_turn.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: ollama_chat, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /v1/messages tool result round trip reaches the model"} +- {id: llm.responses.ollama_chat.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: ollama_chat, capability: basic, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /v1/responses returns text and usage"} +- {id: llm.responses.ollama_chat.basic.stream.works, module: llm, tier: P1, subject_endpoint: responses, route: ollama_chat, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /v1/responses streams text deltas, usage and a terminal event"} +- {id: llm.responses.ollama_chat.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: ollama_chat, capability: tool_use, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /v1/responses returns one addressable get_weather call"} +- {id: llm.responses.ollama_chat.multi_turn.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: ollama_chat, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/chat (native tools) over /v1/responses tool result round trip reaches the model"} +- {id: llm.chat_completions.ollama.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama, capability: basic, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /chat/completions returns text and usage"} +- {id: llm.chat_completions.ollama.basic.stream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /chat/completions streams text deltas, usage and a terminal event"} +- {id: llm.chat_completions.ollama.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama, capability: tool_use, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /chat/completions returns one addressable get_weather call"} +- {id: llm.chat_completions.ollama.multi_turn.nonstream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /chat/completions tool result round trip reaches the model"} +- {id: llm.chat_completions.ollama.tool_use.stream.works, module: llm, tier: P1, subject_endpoint: chat_completions, route: ollama, capability: tool_use, streaming: stream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) streams the call as tool_call deltas ending in finish_reason tool_calls, the shape coding agents like OpenCode execute; prompt-based JSON used to arrive as plain text with finish_reason stop (GitHub issue #35711)"} +- {id: llm.messages.ollama.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: ollama, capability: basic, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /v1/messages returns text and usage"} +- {id: llm.messages.ollama.basic.stream.works, module: llm, tier: P1, subject_endpoint: messages, route: ollama, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /v1/messages streams text deltas, usage and a terminal event"} +- {id: llm.messages.ollama.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: ollama, capability: tool_use, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /v1/messages returns one addressable get_weather call"} +- {id: llm.messages.ollama.multi_turn.nonstream.works, module: llm, tier: P1, subject_endpoint: messages, route: ollama, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /v1/messages tool result round trip reaches the model"} +- {id: llm.responses.ollama.basic.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: ollama, capability: basic, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /v1/responses returns text and usage"} +- {id: llm.responses.ollama.basic.stream.works, module: llm, tier: P1, subject_endpoint: responses, route: ollama, capability: basic, streaming: stream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /v1/responses streams text deltas, usage and a terminal event"} +- {id: llm.responses.ollama.tool_use.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: ollama, capability: tool_use, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /v1/responses returns one addressable get_weather call"} +- {id: llm.responses.ollama.multi_turn.nonstream.works, module: llm, tier: P1, subject_endpoint: responses, route: ollama, capability: multi_turn, streaming: nonstream, assertions: [works], source: "llm_translation/test_ollama_e2e.py", rationale: "Ollama /api/generate (prompt-based tools) over /v1/responses tool result round trip reaches the model"} - {id: llm.chat_completions.sail.service_tier.nonstream.cost_logged, module: llm, tier: P1, subject_endpoint: chat_completions, route: sail, capability: service_tier, streaming: nonstream, assertions: [cost_logged], source: "llm_translation/test_sail_e2e.py", rationale: "service_tier balanced and auto map to Sail completion windows and bill the matching price columns; Sail serves flex only to background responses and Batch"} - {id: llm.chat_completions.sail.service_tier.nonstream.rejects_unknown_tier, module: llm, tier: P1, subject_endpoint: chat_completions, route: sail, capability: service_tier, streaming: nonstream, assertions: [rejects_unknown_tier], source: "llm_translation/test_sail_e2e.py", rationale: "A service_tier Sail has no completion window for is a 400 without drop_params"} - {id: llm.chat_completions.sail.service_tier.nonstream.drops_unknown_tier_and_bills_asap, module: llm, tier: P1, subject_endpoint: chat_completions, route: sail, capability: service_tier, streaming: nonstream, assertions: [drops_unknown_tier_and_bills_asap], source: "llm_translation/test_sail_e2e.py", rationale: "An unknown service_tier under drop_params is dropped and billed at asap in both the cost header and spend log"} diff --git a/tests/e2e/coverage_registry/schema.py b/tests/e2e/coverage_registry/schema.py index d089b9c1ed8..931085cbbc2 100644 --- a/tests/e2e/coverage_registry/schema.py +++ b/tests/e2e/coverage_registry/schema.py @@ -55,6 +55,8 @@ LlmRoute = Literal[ "cohere", "gemini", "hosted_vllm", + "ollama", + "ollama_chat", "openai", "sail", "together_ai", diff --git a/tests/e2e/llm_translation/test_ollama_e2e.py b/tests/e2e/llm_translation/test_ollama_e2e.py new file mode 100644 index 00000000000..6dea5a51562 --- /dev/null +++ b/tests/e2e/llm_translation/test_ollama_e2e.py @@ -0,0 +1,258 @@ +"""Ollama behind the proxy on /chat/completions, /v1/messages and /v1/responses. + +Ollama has two litellm routes with different tool plumbing: `ollama_chat/` calls +/api/chat and forwards native tools, while `ollama/` calls /api/generate, which +has no tools field, so litellm prompts the model for a JSON function call and +turns that JSON back into a tool call. Each route runs the same conversation +contract as the conversational matrix on every surface, through the matrix's +SDK-backed surfaces, plus a streamed tool call on chat completions, the shape +coding agents such as OpenCode consume. + +The deployments set drop_params because Ollama has no parallel_tool_calls, +which the matrix surfaces send alongside a forced tool_choice. Live only: Ollama +Cloud has no provider edge mount, and its requests are not priced in the cost +map, so there is no cost cell here. +""" + +from __future__ import annotations + +import json +from collections.abc import Iterator, Mapping +from dataclasses import dataclass +from itertools import chain, product +from types import MappingProxyType +from typing import Final, Literal, cast + +import pytest +from _pytest.mark.structures import ParameterSet +from e2e_config import unique_marker +from e2e_metadata import Capability as SubjectCapability +from e2e_metadata import Domain, Mode, Provider, Route, Subject, meta +from lifecycle import ResourceManager +from llm_translation.conversational_matrix import ( + GREETING_PROMPT, + INSTRUCTIONS, + MAX_OUTPUT_TOKENS, + SURFACES, + WEATHER_PROMPT, + WEATHER_REPORT, + WEATHER_TOOL_DESCRIPTION, + WEATHER_TOOL_NAME, + WEATHER_TOOL_SCHEMA, + Surface, + SurfaceName, + ToolCall, + WeatherArgs, + build_surfaces, +) +from llm_translation.sdk_clients import NO_PROXY_CACHE, SdkClients +from models import LiteLLMParamsBody +from openai.types.chat import ChatCompletionChunk, ChatCompletionToolParam +from proxy_client import ProxyClient + +pytestmark = pytest.mark.e2e + +OllamaRoute = Literal["ollama_chat", "ollama"] +Capability = Literal["basic", "tool_use", "multi_turn"] +Streaming = Literal["stream", "nonstream"] + +OLLAMA_API_BASE: Final = "https://ollama.com" +OLLAMA_MODEL: Final = "gemma4:31b" +ROUTES: Final[tuple[OllamaRoute, ...]] = ("ollama_chat", "ollama") + + +@dataclass(frozen=True, slots=True) +class Cell: + surface: SurfaceName + route: OllamaRoute + + @property + def id(self) -> str: + return f"{self.surface}-{self.route}" + + +def _cells(capability: Capability, streaming: Streaming) -> tuple[ParameterSet, ...]: + return tuple( + pytest.param( + Cell(surface=surface, route=route), + id=f"{surface}-{route}", + marks=pytest.mark.covers(f"llm.{surface}.{route}.{capability}.{streaming}.works"), + ) + for surface, route in product(SURFACES, ROUTES) + ) + + +def _streamed_tool_cells() -> tuple[ParameterSet, ...]: + return tuple( + pytest.param(route, id=route, marks=pytest.mark.covers(f"llm.chat_completions.{route}.tool_use.stream.works")) + for route in ROUTES + ) + + +def _register(proxy: ProxyClient, resources: ResourceManager, route: OllamaRoute) -> str: + alias: Final = f"e2e-ollama-{route}-{unique_marker()}" + model_id: Final = proxy.create_model( + alias, + LiteLLMParamsBody( + model=f"{route}/{OLLAMA_MODEL}", + api_base=OLLAMA_API_BASE, + api_key="os.environ/OLLAMA_API_KEY", + drop_params=True, + ), + ) + resources.defer(lambda: proxy.delete_model(model_id)) + return alias + + +@pytest.fixture(scope="module") +def aliases(proxy: ProxyClient) -> Iterator[Mapping[OllamaRoute, str]]: + resources: Final = ResourceManager(client=proxy) + try: + yield MappingProxyType({route: _register(proxy, resources, route) for route in ROUTES}) + finally: + resources.teardown() + + +@pytest.fixture(scope="module") +def surfaces(sdk: SdkClients) -> Mapping[SurfaceName, Surface]: + return build_surfaces(sdk) + + +def _weather_call(cell: Cell, surface: Surface, key: str, model: str) -> ToolCall: + reply: Final = surface.reply(key, model, WEATHER_PROMPT, with_tool=True) + assert len(reply.tool_calls) == 1, ( + f"{cell.id}: expected one {WEATHER_TOOL_NAME} call, got {reply.tool_calls} text={reply.text!r}" + ) + call: Final = reply.tool_calls[0] + assert call.name == WEATHER_TOOL_NAME, f"{cell.id}: called {call.name!r}, not {WEATHER_TOOL_NAME!r}" + assert call.call_id, f"{cell.id}: tool call has no id, so the caller cannot answer it: {call}" + assert "paris" in call.parsed().location.lower(), f"{cell.id}: tool arguments lost the location: {call}" + return call + + +def _weather_tool() -> ChatCompletionToolParam: + return { + "type": "function", + "function": { + "name": WEATHER_TOOL_NAME, + "description": WEATHER_TOOL_DESCRIPTION, + "parameters": dict(WEATHER_TOOL_SCHEMA), + }, + } + + +def _subject(mode: Mode, *, tools: bool, route: Route | None = None) -> Subject: + return Subject( + domain=Domain.LLM_TRANSLATION, + route=route, + providers=(Provider.OLLAMA,), + models=(OLLAMA_MODEL,), + capabilities=(SubjectCapability.FUNCTION_CALLING,) if tools else (), + mode=mode, + ) + + +class TestOllamaConversation: + @pytest.mark.parametrize("cell", _cells("basic", "nonstream")) + @meta(_subject(Mode.NONSTREAM, tools=False)) + def test_reply_carries_assistant_text_and_usage( + self, + cell: Cell, + aliases: Mapping[OllamaRoute, str], + surfaces: Mapping[SurfaceName, Surface], + resources: ResourceManager, + ) -> None: + reply: Final = surfaces[cell.surface].reply(resources.key(), aliases[cell.route], GREETING_PROMPT) + + assert reply.response_id, f"{cell.id}: response has no id" + assert reply.text.strip(), f"{cell.id}: response carried no assistant text" + assert reply.usage is not None and reply.usage.input_tokens > 0 and reply.usage.output_tokens > 0, ( + f"{cell.id}: usage missing or zero: {reply.usage}" + ) + assert reply.call_id_header, f"{cell.id}: x-litellm-call-id header missing" + + @pytest.mark.parametrize("cell", _cells("basic", "stream")) + @meta(_subject(Mode.STREAM, tools=False)) + def test_stream_delivers_text_usage_and_a_terminal_event( + self, + cell: Cell, + aliases: Mapping[OllamaRoute, str], + surfaces: Mapping[SurfaceName, Surface], + resources: ResourceManager, + ) -> None: + streamed: Final = surfaces[cell.surface].stream(resources.key(), aliases[cell.route], GREETING_PROMPT) + + assert streamed.event_count > 1, f"{cell.id}: stream arrived as {streamed.event_count} event(s)" + assert streamed.text.strip(), f"{cell.id}: stream carried no text deltas" + assert streamed.finished, f"{cell.id}: stream never sent its terminal event" + assert streamed.usage_reported, f"{cell.id}: stream never reported usage" + + @pytest.mark.parametrize("cell", _cells("tool_use", "nonstream")) + @meta(_subject(Mode.NONSTREAM, tools=True)) + def test_tool_call_is_returned_named_and_addressable( + self, + cell: Cell, + aliases: Mapping[OllamaRoute, str], + surfaces: Mapping[SurfaceName, Surface], + resources: ResourceManager, + ) -> None: + _ = _weather_call(cell, surfaces[cell.surface], resources.key(), aliases[cell.route]) + + @pytest.mark.parametrize("cell", _cells("multi_turn", "nonstream")) + @meta(_subject(Mode.NONSTREAM, tools=True)) + def test_tool_result_round_trip_reaches_the_model( + self, + cell: Cell, + aliases: Mapping[OllamaRoute, str], + surfaces: Mapping[SurfaceName, Surface], + resources: ResourceManager, + ) -> None: + key: Final = resources.key() + model: Final = aliases[cell.route] + surface: Final = surfaces[cell.surface] + call: Final = _weather_call(cell, surface, key, model) + + answer: Final = surface.reply_to_tool_result(key, model, WEATHER_PROMPT, call, WEATHER_REPORT) + assert "22" in answer.text, f"{cell.id}: the model never saw the tool result: {answer.text!r}" + + +class TestOllamaStreamedToolCall: + @pytest.mark.parametrize("route", _streamed_tool_cells()) + @meta(_subject(Mode.STREAM, tools=True, route=Route.CHAT_COMPLETIONS)) + def test_tool_call_streams_as_tool_call_deltas( + self, + route: OllamaRoute, + aliases: Mapping[OllamaRoute, str], + sdk: SdkClients, + resources: ResourceManager, + ) -> None: + chunks: Final[tuple[ChatCompletionChunk, ...]] = tuple( + sdk.openai(resources.key()).chat.completions.create( + model=aliases[route], + messages=[ + {"role": "system", "content": INSTRUCTIONS}, + {"role": "user", "content": WEATHER_PROMPT}, + ], + tools=[_weather_tool()], + max_completion_tokens=MAX_OUTPUT_TOKENS, + stream=True, + extra_body=NO_PROXY_CACHE, + ) + ) + choices: Final = tuple(chunk.choices[0] for chunk in chunks if chunk.choices) + text: Final = "".join(choice.delta.content or "" for choice in choices) + deltas: Final = tuple(chain.from_iterable(choice.delta.tool_calls or () for choice in choices)) + call_ids: Final = tuple(delta.id for delta in deltas if delta.id) + indexes: Final = frozenset(delta.index for delta in deltas) + functions: Final = tuple(delta.function for delta in deltas if delta.function is not None) + names: Final = tuple(function.name for function in functions if function.name) + arguments: Final = "".join(function.arguments or "" for function in functions) + finish_reasons: Final = tuple(choice.finish_reason for choice in choices if choice.finish_reason is not None) + + assert names == (WEATHER_TOOL_NAME,), f"{route}: streamed tool names {names}, text={text!r}" + assert len(call_ids) == 1, f"{route}: expected one streamed tool call id, got {call_ids}" + assert indexes == {0}, f"{route}: streamed tool call deltas used indexes {sorted(indexes)}" + assert WEATHER_TOOL_NAME not in text, f"{route}: the tool call leaked into assistant text: {text!r}" + location: Final = WeatherArgs.model_validate(cast(object, json.loads(arguments))).location + assert "paris" in location.lower(), f"{route}: streamed tool arguments lost the location: {arguments!r}" + assert finish_reasons[-1:] == ("tool_calls",), f"{route}: stream finished with {finish_reasons}" diff --git a/tests/e2e/models.py b/tests/e2e/models.py index 8796393f269..3ae200ff5e7 100644 --- a/tests/e2e/models.py +++ b/tests/e2e/models.py @@ -1257,6 +1257,7 @@ class LiteLLMParamsBody(BaseModel): api_version: str | None = None realtime_protocol: str | None = None allowed_openai_params: list[str] | None = None + drop_params: bool | None = None aws_access_key_id: str | None = Field(default=None, repr=False) aws_secret_access_key: str | None = Field(default=None, repr=False) aws_region_name: str | None = None diff --git a/tests/integration/providers/test_ollama_prompt_tools_chaos.py b/tests/integration/providers/test_ollama_prompt_tools_chaos.py new file mode 100644 index 00000000000..243c85b2a64 --- /dev/null +++ b/tests/integration/providers/test_ollama_prompt_tools_chaos.py @@ -0,0 +1,404 @@ +import asyncio +import json +import re +import signal +import threading +import uuid +from collections.abc import Callable +from dataclasses import dataclass +from pathlib import Path +from queue import SimpleQueue +from types import MappingProxyType +from typing import Final, Literal +from urllib.parse import urlsplit + +import httpx +import psutil +import pytest +import yaml +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.process import graceful_stop_seconds, owned_proxy_process +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "llama3-prompt-tools-chaos" +_API_KEY: Final = "synthetic-ollama-key" +_CONFIG_MODEL: Final = "ollama-prompt-tools-chaos" +_INSTRUCTION: Final = ( + 'To call a function, reply with JSON ONLY in this format {"name": "function_name", ' + '"arguments":{"argument_name": "argument_value"}}. Once a function result answers the request, ' + "reply to the user in plain text instead of calling a function again. " + "The following functions are available to you:" +) +_CALL_ID: Final = "call_prompt_tools_chaos_1" +_ARGUMENTS: Final[dict[str, JsonValue]] = {"city": "Paris"} +_PARAMETERS: Final[dict[str, JsonValue]] = { + "type": "object", + "properties": {"city": {"type": "string"}}, + "required": ["city"], +} +_CHAT_TOOL: Final[dict[str, JsonValue]] = { + "type": "function", + "function": {"name": "get_weather", "description": "Weather for a city", "parameters": _PARAMETERS}, +} +_ANTHROPIC_TOOL: Final[dict[str, JsonValue]] = { + "name": "get_weather", + "description": "Weather for a city", + "input_schema": _PARAMETERS, +} +_RESPONSES_TOOL: Final[dict[str, JsonValue]] = { + "type": "function", + "name": "get_weather", + "description": "Weather for a city", + "parameters": _PARAMETERS, +} +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_MARKER: Final = re.compile(r"marker-([0-9a-f]{32})") +_STARTED_WORKER: Final = re.compile(r"Started server process \[(\d+)\]") +_OWNED_PROXY_CELL_SECONDS: Final = 2 * graceful_stop_seconds() + 120 + +Endpoint = Literal["chat", "messages", "responses"] + + +@dataclass(frozen=True, slots=True) +class _Call: + endpoint: Endpoint + stream: bool + marker: str + + +@dataclass(frozen=True, slots=True) +class _Served: + call: _Call + status: int + text: str + + +def _result(marker: str) -> str: + return f"Paris: 22 degrees Celsius marker-{marker}" + + +def _answer(marker: str) -> str: + return f"answer marker-{marker}" + + +def _path(endpoint: Endpoint) -> str: + match endpoint: + case "chat": + return "/v1/chat/completions" + case "messages": + return "/v1/messages" + case "responses": + return "/v1/responses" + + +def _body(model: str, call: _Call) -> dict[str, JsonValue]: + question: Final = "What is the weather in Paris?" + common: Final[dict[str, JsonValue]] = { + "model": model, + "stream": call.stream, + "num_retries": 0, + "cache": {"no-cache": True}, + } + match call.endpoint: + case "chat": + return { + **common, + "tools": [_CHAT_TOOL], + "messages": [ + {"role": "user", "content": question}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": _CALL_ID, + "type": "function", + "function": {"name": "get_weather", "arguments": json.dumps(_ARGUMENTS)}, + } + ], + }, + {"role": "tool", "tool_call_id": _CALL_ID, "content": _result(call.marker)}, + ], + } + case "messages": + return { + **common, + "max_tokens": 64, + "tools": [_ANTHROPIC_TOOL], + "messages": [ + {"role": "user", "content": question}, + { + "role": "assistant", + "content": [{"type": "tool_use", "id": _CALL_ID, "name": "get_weather", "input": _ARGUMENTS}], + }, + { + "role": "user", + "content": [{"type": "tool_result", "tool_use_id": _CALL_ID, "content": _result(call.marker)}], + }, + ], + } + case "responses": + return { + **common, + "store": False, + "tools": [_RESPONSES_TOOL], + "input": [ + {"role": "user", "content": question}, + { + "type": "function_call", + "call_id": _CALL_ID, + "name": "get_weather", + "arguments": json.dumps(_ARGUMENTS), + }, + {"type": "function_call_output", "call_id": _CALL_ID, "output": _result(call.marker)}, + ], + } + + +def _generate_reply(marker: str, stream: bool, drop_connection: bool = False) -> Reply: + done: Final = { + "model": _BACKEND, + "created_at": "2026-10-07T00:00:00Z", + "response": "", + "done": True, + "done_reason": "stop", + "prompt_eval_count": 30, + "eval_count": 12, + } + if not stream: + return Reply(body=json.dumps({**done, "response": _answer(marker)}).encode(), drop_connection=drop_connection) + pieces: Final = ("answer ", f"marker-{marker}") + frames: Final = ( + *( + {"model": _BACKEND, "created_at": "2026-10-07T00:00:00Z", "response": piece, "done": False} + for piece in pieces + ), + done, + ) + return Reply( + content_type="application/x-ndjson", + chunks=tuple(json.dumps(frame).encode() + b"\n" for frame in frames), + drop_connection=drop_connection, + ) + + +def _is_generate(request: Request) -> bool: + return (request.method, request.target) == ("POST", "/api/generate") + + +def _marker_of(request: Request) -> str: + found: Final = _MARKER.search(request.body.decode()) + assert found is not None, request.body + return found.group(1) + + +def _echo(request: Request) -> Reply: + if not _is_generate(request): + return Reply(body=b"{}") + stream: Final = _JSON_OBJECT.validate_json(request.body).get("stream") is True + return _generate_reply(_marker_of(request), stream) + + +def _assert_each_prompt_is_instructed_once(received: tuple[Request, ...], markers: frozenset[str]) -> None: + generates: Final = tuple(request for request in received if _is_generate(request)) + assert sorted(_marker_of(request) for request in generates) == sorted(markers) + for request in generates: + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["format"] == "json", sorted(body) + prompt: Final = body["prompt"] + assert isinstance(prompt, str) + assert prompt.count(_INSTRUCTION) == 1, prompt + assert set(_MARKER.findall(prompt)) == {_marker_of(request)}, prompt + + +def _response_id(served: _Served) -> str | None: + if served.call.endpoint == "responses": + return None + if not served.call.stream: + identity: Final = _JSON_OBJECT.validate_json(served.text)["id"] + assert isinstance(identity, str) + return identity + for line in served.text.splitlines(): + if not line.startswith("data: ") or line == "data: [DONE]": + continue + payload: Final = _JSON_OBJECT.validate_json(line.removeprefix("data: ")) + if served.call.endpoint == "chat": + first: Final = payload["id"] + assert isinstance(first, str) + return first + if payload.get("type") == "message_start": + message: Final = payload["message"] + assert isinstance(message, dict) and isinstance(message["id"], str) + return message["id"] + raise AssertionError(served.text) + + +def _spend_statuses(model: str, expected: int) -> MappingProxyType[str, str]: + rows: Final = eventually( + lambda: read_rows('SELECT request_id, status FROM "LiteLLM_SpendLogs" WHERE model_group=%s', (model,)), + lambda found: len(found) >= expected, + seconds=70, + ) + statuses: Final = MappingProxyType({str(row["request_id"]): str(row["status"]) for row in rows}) + assert len(statuses) == len(rows) == expected, rows + return statuses + + +def _successes(statuses: MappingProxyType[str, str]) -> frozenset[str]: + return frozenset(identity for identity, status in statuses.items() if status == "success") + + +async def _send(client: httpx.AsyncClient, key: str, model: str, call: _Call) -> _Served: + async with client.stream( + "POST", + _path(call.endpoint), + json=_body(model, call), + headers={"Authorization": f"Bearer {key}", "anthropic-version": "2023-06-01"}, + ) as response: + raw: Final = await response.aread() + return _Served(call=call, status=response.status_code, text=raw.decode()) + + +async def _burst( + base_url: str, key: str, model: str, calls: tuple[_Call, ...], *, tolerate_transport_errors: bool = False +) -> tuple[_Served, ...]: + async with httpx.AsyncClient(base_url=base_url, timeout=60, trust_env=False) as client: + results: Final = await asyncio.gather( + *(_send(client, key, model, call) for call in calls), return_exceptions=tolerate_transport_errors + ) + for result in results: + assert not isinstance(result, BaseException) or isinstance(result, httpx.TransportError), repr(result) + return tuple(result for result in results if isinstance(result, _Served)) + + +def _calls(count: int, endpoints: tuple[Endpoint, ...], stream: Callable[[int], bool]) -> tuple[_Call, ...]: + return tuple( + _Call(endpoint=endpoints[index % len(endpoints)], stream=stream(index), marker=uuid.uuid4().hex) + for index in range(count) + ) + + +def _assert_answered_with_its_own_marker(served: _Served) -> None: + assert served.status == 200, served.text + assert set(_MARKER.findall(served.text)) == {served.call.marker}, served.text + + +async def test_concurrent_tool_result_turns_across_endpoints_each_get_their_own_instructed_prompt( + gateway: Gateway, +) -> None: + calls: Final = _calls(24, ("chat", "messages", "responses"), lambda index: index % 2 == 0) + with wire_server(_echo) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + served: Final = await _burst(str(gateway.client.base_url), gateway.key, model, calls) + assert len(served) == 24 + for item in served: + _assert_answered_with_its_own_marker(item) + _assert_each_prompt_is_instructed_once(wire.drain(), frozenset(call.marker for call in calls)) + known: Final = frozenset(identity for identity in map(_response_id, served) if identity is not None) + assert len(known) == 16, known + statuses: Final = _spend_statuses(model, 24) + assert _successes(statuses) == frozenset(statuses), statuses + assert known <= _successes(statuses) + + +async def test_dropped_ollama_connections_fail_their_callers_and_the_rest_keep_their_prompts(gateway: Gateway) -> None: + calls: Final = _calls(12, ("chat",), lambda index: index % 2 == 1) + dropped: Final = frozenset(call.marker for index, call in enumerate(calls) if index % 3 == 0) + + def respond(request: Request) -> Reply: + if not _is_generate(request): + return Reply(body=b"{}") + marker: Final = _marker_of(request) + stream: Final = _JSON_OBJECT.validate_json(request.body).get("stream") is True + return _generate_reply(marker, stream, drop_connection=marker in dropped) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + served: Final = await _burst(str(gateway.client.base_url), gateway.key, model, calls) + assert len(served) == 12 + for item in served: + if item.call.marker in dropped: + assert item.status == 500, item.text + assert "APIConnectionError" in item.text and "marker-" not in item.text, item.text + else: + _assert_answered_with_its_own_marker(item) + recovery: Final = _Call(endpoint="chat", stream=False, marker=uuid.uuid4().hex) + (recovered,) = await _burst(str(gateway.client.base_url), gateway.key, model, (recovery,)) + _assert_answered_with_its_own_marker(recovered) + _assert_each_prompt_is_instructed_once(wire.drain(), frozenset(call.marker for call in (*calls, recovery))) + answered: Final = tuple(item for item in served if item.call.marker not in dropped) + survivors: Final = frozenset( + identity for identity in map(_response_id, (*answered, recovered)) if identity is not None + ) + assert len(survivors) == 9, survivors + statuses: Final = _spend_statuses(model, 13) + assert _successes(statuses) == survivors, statuses + assert sum(status == "failure" for status in statuses.values()) == len(dropped), statuses + + +def _chaos_config(wire: Wire, tmp_path: Path) -> Path: + base: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + deployment: Final = { + "model_name": _CONFIG_MODEL, + "litellm_params": {"model": f"ollama/{_BACKEND}", "api_base": wire.url, "api_key": _API_KEY}, + } + path: Final = tmp_path / "ollama-prompt-tools-chaos.yaml" + path.write_text(yaml.safe_dump({**base, "model_list": [deployment]})) + return path + + +def _open_upstream_connections(pid: int, upstream: str) -> int: + port: Final = urlsplit(upstream).port + return sum( + 1 + for connection in psutil.Process(pid).net_connections(kind="tcp") + if connection.status == psutil.CONN_ESTABLISHED and connection.raddr and connection.raddr.port == port + ) + + +@pytest.mark.timeout(_OWNED_PROXY_CELL_SECONDS) +async def test_worker_sigkill_mid_burst_leaves_the_sibling_instructing_ollama(gateway: Gateway, tmp_path: Path) -> None: + calls: Final = _calls(20, ("chat",), lambda _: False) + release: Final = threading.Event() + held_markers: Final[SimpleQueue[str]] = SimpleQueue() + + def held(request: Request) -> Reply: + if not _is_generate(request): + return Reply(body=b"{}") + held_markers.put(_marker_of(request)) + assert release.wait(timeout=60), "The burst was never released" + return _echo(request) + + with wire_server(held) as wire: + path: Final = _chaos_config(wire, tmp_path) + with owned_proxy_process(gateway, tmp_path, {}, config=path, workers=2) as owned: + candidate: Final = owned.gateway + workers: Final = eventually( + lambda: tuple(int(pid) for pid in _STARTED_WORKER.findall(owned.log.read_text())), + lambda pids: len(pids) == 2, + seconds=30, + ) + burst: Final = asyncio.create_task( + _burst( + str(candidate.client.base_url), candidate.key, _CONFIG_MODEL, calls, tolerate_transport_errors=True + ) + ) + await asyncio.to_thread(eventually, held_markers.qsize, lambda size: size == 20, 60) + held_by: Final = MappingProxyType({pid: _open_upstream_connections(pid, wire.url) for pid in workers}) + assert sum(held_by.values()) == 20, held_by + victim_pid, survivor_pid = sorted(workers, key=held_by.__getitem__) + victim: Final = psutil.Process(victim_pid) + victim.suspend() + victim.send_signal(signal.SIGKILL) + release.set() + served: Final = await burst + assert held_by[survivor_pid] >= 10, held_by + assert len(served) == held_by[survivor_pid], (held_by, len(served)) + for item in served: + _assert_answered_with_its_own_marker(item) + follow_up: Final = _Call(endpoint="chat", stream=False, marker=uuid.uuid4().hex) + (answered,) = await _burst(str(candidate.client.base_url), candidate.key, _CONFIG_MODEL, (follow_up,)) + _assert_answered_with_its_own_marker(answered) + _assert_each_prompt_is_instructed_once(wire.drain(), frozenset(call.marker for call in (*calls, follow_up))) diff --git a/tests/integration/providers/test_ollama_prompt_tools_wire.py b/tests/integration/providers/test_ollama_prompt_tools_wire.py new file mode 100644 index 00000000000..115fdc65e33 --- /dev/null +++ b/tests/integration/providers/test_ollama_prompt_tools_wire.py @@ -0,0 +1,796 @@ +import itertools +import json +import uuid +from collections.abc import Callable, Iterator, Sequence +from contextlib import contextmanager +from typing import Final + +import anthropic +import openai +import pytest +from openai.types.chat import ChatCompletionChunk +from openai.types.chat.chat_completion_chunk import Choice as ChunkChoice +from openai.types.chat.chat_completion_chunk import ChoiceDeltaToolCall +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.wire import Reply, Request, Wire, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "llama3-prompt-tools" +_API_KEY: Final = "synthetic-ollama-key" +_INSTRUCTION: Final = ( + 'To call a function, reply with JSON ONLY in this format {"name": "function_name", ' + '"arguments":{"argument_name": "argument_value"}}. Once a function result answers the request, ' + "reply to the user in plain text instead of calling a function again. " + "The following functions are available to you:" +) +_QUESTION: Final = "What is the weather in Paris?" +_RESULT: Final = "Paris: 22 degrees Celsius, clear skies" +_ANSWER: Final = "Paris is 22 degrees Celsius with clear skies." +_CALL_ID: Final = "call_prompt_tools_1" +_ARGUMENTS: Final[dict[str, JsonValue]] = {"city": "Paris"} +_CALL_JSON: Final = json.dumps({"name": "get_weather", "arguments": _ARGUMENTS}) +_CALL_JSON_FIELDS: Final = frozenset({"get_weather"}) +_PARAMETERS: Final[dict[str, JsonValue]] = { + "type": "object", + "properties": {"city": {"type": "string"}}, + "required": ["city"], +} +_WEATHER_TOOL: Final[dict[str, JsonValue]] = { + "type": "function", + "function": {"name": "get_weather", "description": "Weather for a city", "parameters": _PARAMETERS}, +} +_TIME_TOOL: Final[dict[str, JsonValue]] = { + "type": "function", + "function": {"name": "get_time", "description": "Local time for a city", "parameters": _PARAMETERS}, +} +_ANTHROPIC_TOOL: Final[dict[str, JsonValue]] = { + "name": "get_weather", + "description": "Weather for a city", + "input_schema": _PARAMETERS, +} +_RESPONSES_TOOL: Final[dict[str, JsonValue]] = { + "type": "function", + "name": "get_weather", + "description": "Weather for a city", + "parameters": _PARAMETERS, +} +_NO_CACHE: Final[dict[str, JsonValue]] = {"no-cache": True} +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_RAW_ANTHROPIC_EVENTS: Final = frozenset( + { + "message_start", + "content_block_start", + "content_block_delta", + "content_block_stop", + "message_delta", + "message_stop", + } +) + + +def _generate_reply(text: str) -> Reply: + return Reply( + body=json.dumps( + { + "model": _BACKEND, + "created_at": "2026-10-07T00:00:00Z", + "response": text, + "done": True, + "done_reason": "stop", + "prompt_eval_count": 30, + "eval_count": 12, + } + ).encode() + ) + + +def _streamed_reply(text: str) -> Reply: + pieces: Final = tuple(text[index : index + 7] for index in range(0, len(text), 7)) + lines: Final = tuple( + json.dumps({"model": _BACKEND, "created_at": "2026-10-07T00:00:00Z", "response": piece, "done": False}).encode() + + b"\n" + for piece in pieces + ) + final: Final = ( + json.dumps( + { + "model": _BACKEND, + "created_at": "2026-10-07T00:00:00Z", + "response": "", + "done": True, + "done_reason": "stop", + "prompt_eval_count": 30, + "eval_count": 12, + } + ).encode() + + b"\n" + ) + return Reply(content_type="application/x-ndjson", chunks=(*lines, final)) + + +def _is_generate(request: Request) -> bool: + return (request.method, request.target) == ("POST", "/api/generate") + + +@contextmanager +def _ollama_server(respond: Callable[[Request], Reply]) -> Iterator[Wire]: + with wire_server(lambda request: respond(request) if _is_generate(request) else Reply(body=b"{}")) as wire: + yield wire + + +def _generate_calls(wire: Wire) -> tuple[Request, ...]: + return tuple(request for request in wire.drain() if _is_generate(request)) + + +def _only_generate(wire: Wire) -> dict[str, JsonValue]: + received: Final = _generate_calls(wire) + assert len(received) == 1, [(request.method, request.target) for request in received] + assert received[0].headers["authorization"] == f"Bearer {_API_KEY}" + return _JSON_OBJECT.validate_json(received[0].body) + + +def _prompt_of(body: dict[str, JsonValue]) -> str: + assert body["model"] == _BACKEND + assert body["format"] == "json" + assert "tools" not in body and "messages" not in body, sorted(body) + prompt: Final = body["prompt"] + assert isinstance(prompt, str) + return prompt + + +def _assert_instructed_once(prompt: str, *tool_names: str) -> None: + assert prompt.count(_INSTRUCTION) == 1, prompt + assert prompt.count("### System:") == 1, prompt + for name in tool_names: + assert f"'name': '{name}'" in prompt, prompt + + +def _assert_tool_turn(prompt: str, result: str = _RESULT) -> None: + _assert_instructed_once(prompt, "get_weather") + assert f"### User:\n{_QUESTION}\n\n### Assistant:\n{_CALL_JSON}\n\n### User:\n{result}\n\n" in prompt, prompt + + +def _spend_row(identity: str) -> dict[str, JsonValue]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT model_group, status, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE request_id=%s', + (identity,), + ), + lambda found: len(found) == 1, + seconds=70, + ) + return rows[0] + + +def _model_spend_rows(model: str, expected: int) -> list[dict[str, JsonValue]]: + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, status, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" WHERE model_group=%s', + (model,), + ), + lambda found: len(found) >= expected, + seconds=70, + ) + assert len({row["request_id"] for row in rows}) == len(rows) == expected, rows + return rows + + +def _billed(model: str) -> dict[str, JsonValue]: + return {"model_group": model, "status": "success", "prompt_tokens": 30, "completion_tokens": 12} + + +def _openai_client(gateway: Gateway) -> openai.OpenAI: + return openai.OpenAI(base_url=str(gateway.client.base_url) + "/v1", api_key=gateway.key, max_retries=0) + + +def _async_openai_client(gateway: Gateway) -> openai.AsyncOpenAI: + return openai.AsyncOpenAI(base_url=str(gateway.client.base_url) + "/v1", api_key=gateway.key, max_retries=0) + + +def _anthropic_client(gateway: Gateway) -> anthropic.Anthropic: + return anthropic.Anthropic(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0) + + +def _async_anthropic_client(gateway: Gateway) -> anthropic.AsyncAnthropic: + return anthropic.AsyncAnthropic(base_url=str(gateway.client.base_url), api_key=gateway.key, max_retries=0) + + +def _stream_choices(chunks: Sequence[ChatCompletionChunk]) -> Iterator[ChunkChoice]: + for chunk in chunks: + yield from chunk.choices + + +def _delta_tool_calls(choices: Sequence[ChunkChoice]) -> Iterator[ChoiceDeltaToolCall]: + for choice in choices: + yield from choice.delta.tool_calls or () + + +def _first_turn() -> list[dict[str, JsonValue]]: + return [{"role": "user", "content": _QUESTION}] + + +def _second_turn(result: JsonValue = _RESULT) -> list[dict[str, JsonValue]]: + return [ + {"role": "user", "content": _QUESTION}, + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": _CALL_ID, + "type": "function", + "function": {"name": "get_weather", "arguments": json.dumps(_ARGUMENTS)}, + } + ], + }, + {"role": "tool", "tool_call_id": _CALL_ID, "content": result}, + ] + + +def _anthropic_second_turn() -> list[dict[str, JsonValue]]: + return [ + {"role": "user", "content": _QUESTION}, + { + "role": "assistant", + "content": [{"type": "tool_use", "id": _CALL_ID, "name": "get_weather", "input": _ARGUMENTS}], + }, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": _CALL_ID, "content": _RESULT}]}, + ] + + +def _responses_second_turn() -> list[dict[str, JsonValue]]: + return [ + {"role": "user", "content": _QUESTION}, + {"type": "function_call", "call_id": _CALL_ID, "name": "get_weather", "arguments": json.dumps(_ARGUMENTS)}, + {"type": "function_call_output", "call_id": _CALL_ID, "output": _RESULT}, + ] + + +def _post(gateway: Gateway, path: str, body: dict[str, JsonValue], key: str | None = None) -> tuple[int, str]: + response: Final = gateway.request("POST", path, {**body, "cache": _NO_CACHE}, key=key) + return response.status_code, response.text + + +def _post_chat( + gateway: Gateway, model: str, messages: Sequence[dict[str, JsonValue]], **extra: JsonValue +) -> dict[str, JsonValue]: + code, text = _post(gateway, "/v1/chat/completions", {"model": model, "messages": list(messages), **extra}) + assert code == 200, text + return _JSON_OBJECT.validate_json(text) + + +def test_openai_sdk_tool_request_reaches_ollama_as_an_instructed_prompt(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + completion: Final = _openai_client(gateway).chat.completions.create( + model=model, + messages=_first_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_WEATHER_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) + choice: Final = completion.choices[0] + assert choice.finish_reason == "tool_calls" + assert choice.message.tool_calls is not None and len(choice.message.tool_calls) == 1 + call: Final = choice.message.tool_calls[0] + assert call.type == "function" + assert call.function.name == "get_weather" + assert json.loads(call.function.arguments) == _ARGUMENTS + assert completion.usage is not None + assert (completion.usage.prompt_tokens, completion.usage.completion_tokens) == (30, 12) + body: Final = _only_generate(wire) + assert body["stream"] is False + prompt: Final = _prompt_of(body) + _assert_instructed_once(prompt, "get_weather") + assert f"### User:\n{_QUESTION}\n\n" in prompt, prompt + assert "Weather for a city" in prompt, prompt + assert _spend_row(completion.id) == _billed(model) + + +def test_openai_sdk_tool_result_turn_gets_a_plain_text_answer(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + completion: Final = _openai_client(gateway).chat.completions.create( + model=model, + messages=_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_WEATHER_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) + choice: Final = completion.choices[0] + assert choice.finish_reason == "stop" + assert choice.message.content == _ANSWER + assert choice.message.tool_calls is None + _assert_tool_turn(_prompt_of(_only_generate(wire))) + assert _spend_row(completion.id) == _billed(model) + + +async def test_async_openai_sdk_stream_flushes_the_held_tool_call_once(gateway: Gateway) -> None: + with _ollama_server(lambda _: _streamed_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + stream: Final = await _async_openai_client(gateway).chat.completions.create( + model=model, + messages=_first_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_WEATHER_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + stream=True, + stream_options={"include_usage": True}, + extra_body={"cache": _NO_CACHE}, + ) + chunks: Final = [chunk async for chunk in stream] + assert {chunk.id for chunk in chunks} == {chunks[0].id} + choices: Final = tuple(_stream_choices(chunks)) + deltas: Final = tuple(_delta_tool_calls(choices)) + assert len(deltas) == 1, deltas + assert deltas[0].function is not None and deltas[0].function.name == "get_weather" + assert json.loads(deltas[0].function.arguments or "") == _ARGUMENTS + assert [choice.finish_reason for choice in choices if choice.finish_reason] == ["tool_calls"] + usages: Final = [chunk.usage for chunk in chunks if chunk.usage is not None] + assert [(usage.prompt_tokens, usage.completion_tokens) for usage in usages] == [(30, 12)] + body: Final = _only_generate(wire) + assert body["stream"] is True + _assert_instructed_once(_prompt_of(body), "get_weather") + assert _spend_row(chunks[0].id) == _billed(model) + + +async def test_async_openai_sdk_stream_answers_the_tool_result_in_plain_text(gateway: Gateway) -> None: + with _ollama_server(lambda _: _streamed_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + stream: Final = await _async_openai_client(gateway).chat.completions.create( + model=model, + messages=_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_WEATHER_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + stream=True, + stream_options={"include_usage": True}, + extra_body={"cache": _NO_CACHE}, + ) + chunks: Final = [chunk async for chunk in stream] + choices: Final = tuple(_stream_choices(chunks)) + assert "".join(choice.delta.content or "" for choice in choices) == _ANSWER + assert tuple(_delta_tool_calls(choices)) == () + assert [choice.finish_reason for choice in choices if choice.finish_reason] == ["stop"] + body: Final = _only_generate(wire) + assert body["stream"] is True + _assert_tool_turn(_prompt_of(body)) + assert _spend_row(chunks[0].id) == _billed(model) + + +def test_anthropic_sdk_tool_request_comes_back_as_tool_use(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + message: Final = _anthropic_client(gateway).messages.create( + model=model, + max_tokens=64, + messages=_first_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_ANTHROPIC_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) + assert message.stop_reason == "tool_use" + assert [block.type for block in message.content] == ["tool_use"] + block: Final = message.content[0] + assert block.type == "tool_use" + assert block.name == "get_weather" + assert block.input == _ARGUMENTS + assert (message.usage.input_tokens, message.usage.output_tokens) == (30, 12) + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + assert f"### User:\n{_QUESTION}\n\n" in prompt, prompt + assert _spend_row(message.id) == _billed(model) + + +def test_anthropic_sdk_tool_result_turn_ends_with_text(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + message: Final = _anthropic_client(gateway).messages.create( + model=model, + max_tokens=64, + messages=_anthropic_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_ANTHROPIC_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) + assert message.stop_reason == "end_turn" + assert [(block.type, getattr(block, "text", None)) for block in message.content] == [("text", _ANSWER)] + _assert_tool_turn(_prompt_of(_only_generate(wire))) + assert _spend_row(message.id) == _billed(model) + + +async def test_async_anthropic_sdk_stream_emits_the_tool_use_block(gateway: Gateway) -> None: + with _ollama_server(lambda _: _streamed_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + async with _async_anthropic_client(gateway).messages.stream( + model=model, + max_tokens=64, + messages=_first_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_ANTHROPIC_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) as stream: + events: Final = [event async for event in stream if event.type in _RAW_ANTHROPIC_EVENTS] + final: Final = await stream.get_final_message() + starts: Final = [event for event in events if event.type == "content_block_start"] + assert [event.content_block.type for event in starts] == ["tool_use"], [event.type for event in events] + assert any( + event.type == "content_block_delta" and event.delta.type == "input_json_delta" for event in events + ), [event.type for event in events] + assert final.stop_reason == "tool_use" + block: Final = final.content[0] + assert block.type == "tool_use" and block.name == "get_weather" and block.input == _ARGUMENTS + body: Final = _only_generate(wire) + assert body["stream"] is True + _assert_instructed_once(_prompt_of(body), "get_weather") + assert _spend_row(final.id) == _billed(model) + + +async def test_async_anthropic_sdk_stream_answers_the_tool_result_in_text(gateway: Gateway) -> None: + with _ollama_server(lambda _: _streamed_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + async with _async_anthropic_client(gateway).messages.stream( + model=model, + max_tokens=64, + messages=_anthropic_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_ANTHROPIC_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) as stream: + texts: Final = [event.text async for event in stream if event.type == "text"] + final: Final = await stream.get_final_message() + assert "".join(texts) == _ANSWER + assert final.stop_reason == "end_turn" + assert [(block.type, getattr(block, "text", None)) for block in final.content] == [("text", _ANSWER)] + body: Final = _only_generate(wire) + assert body["stream"] is True + _assert_tool_turn(_prompt_of(body)) + assert _spend_row(final.id) == _billed(model) + + +def test_openai_sdk_responses_tool_request_comes_back_as_a_function_call(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + response: Final = _openai_client(gateway).responses.create( + model=model, + input=_QUESTION, + tools=[_RESPONSES_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + store=False, + extra_body={"cache": _NO_CACHE}, + ) + assert [item.type for item in response.output] == ["function_call"] + item: Final = response.output[0] + assert item.type == "function_call" + assert item.name == "get_weather" + assert json.loads(item.arguments) == _ARGUMENTS + assert response.usage is not None and (response.usage.input_tokens, response.usage.output_tokens) == (30, 12) + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + assert f"### User:\n{_QUESTION}\n\n" in prompt, prompt + assert _model_spend_rows(model, 1)[0]["status"] == "success" + + +def test_openai_sdk_responses_function_output_turn_gets_a_message(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + response: Final = _openai_client(gateway).responses.create( + model=model, + input=_responses_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON input items + tools=[_RESPONSES_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + store=False, + extra_body={"cache": _NO_CACHE}, + ) + assert [item.type for item in response.output] == ["message"] + assert response.output_text == _ANSWER + _assert_tool_turn(_prompt_of(_only_generate(wire))) + rows: Final = _model_spend_rows(model, 1) + assert (rows[0]["status"], rows[0]["prompt_tokens"], rows[0]["completion_tokens"]) == ("success", 30, 12) + + +async def test_async_openai_sdk_responses_stream_emits_the_function_call_item(gateway: Gateway) -> None: + with _ollama_server(lambda _: _streamed_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + stream: Final = await _async_openai_client(gateway).responses.create( + model=model, + input=_QUESTION, + tools=[_RESPONSES_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + store=False, + stream=True, + extra_body={"cache": _NO_CACHE}, + ) + events: Final = [event async for event in stream] + done_items: Final = [event.item for event in events if event.type == "response.output_item.done"] + assert [item.type for item in done_items] == ["function_call"], [event.type for event in events] + item: Final = done_items[0] + assert item.type == "function_call" and item.name == "get_weather" + assert json.loads(item.arguments) == _ARGUMENTS + completed: Final = [event for event in events if event.type == "response.completed"] + assert len(completed) == 1 + final_calls: Final = [item for item in completed[0].response.output if item.type == "function_call"] + assert [(item.name, json.loads(item.arguments)) for item in final_calls] == [("get_weather", _ARGUMENTS)] + assert completed[0].response.output_text == "" + body: Final = _only_generate(wire) + assert body["stream"] is True + _assert_instructed_once(_prompt_of(body), "get_weather") + assert _model_spend_rows(model, 1)[0]["status"] == "success" + + +async def test_async_openai_sdk_responses_stream_answers_the_function_output_in_text(gateway: Gateway) -> None: + with _ollama_server(lambda _: _streamed_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + stream: Final = await _async_openai_client(gateway).responses.create( + model=model, + input=_responses_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON input items + tools=[_RESPONSES_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + store=False, + stream=True, + extra_body={"cache": _NO_CACHE}, + ) + events: Final = [event async for event in stream] + assert "".join(event.delta for event in events if event.type == "response.output_text.delta") == _ANSWER + completed: Final = [event for event in events if event.type == "response.completed"] + assert len(completed) == 1 and completed[0].response.output_text == _ANSWER + done_types: Final = [event.type for event in events if event.type == "response.output_item.done"] + assert done_types == ["response.output_item.done"] + body: Final = _only_generate(wire) + assert body["stream"] is True + _assert_tool_turn(_prompt_of(body)) + assert _model_spend_rows(model, 1)[0]["status"] == "success" + + +def test_legacy_functions_param_is_instructed_the_same_way(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + function: Final = _WEATHER_TOOL["function"] + payload: Final = _post_chat(gateway, model, _first_turn(), functions=[function]) + choices: Final = payload["choices"] + assert isinstance(choices, list) and len(choices) == 1 + assert _CALL_JSON_FIELDS <= set(json.dumps(choices[0]).split('"')), choices + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + assert "Weather for a city" in prompt, prompt + identity: Final = payload["id"] + assert isinstance(identity, str) + assert _spend_row(identity) == _billed(model) + + +def test_string_system_message_keeps_one_system_section_with_the_instruction(gateway: Gateway) -> None: + system: Final = f"You are a terse weather bot {uuid.uuid4().hex}." + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + _post_chat(gateway, model, [{"role": "system", "content": system}, *_first_turn()], tools=[_WEATHER_TOOL]) + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + assert f"### System:\n{system} {_INSTRUCTION}\n" in prompt, prompt + + +def test_list_system_message_keeps_one_system_section_with_the_instruction(gateway: Gateway) -> None: + system: Final = f"You are a terse weather bot {uuid.uuid4().hex}." + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + _post_chat( + gateway, + model, + [{"role": "system", "content": [{"type": "text", "text": system}]}, *_first_turn()], + tools=[_WEATHER_TOOL], + ) + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + section: Final = prompt.split("### System:\n", 1)[1] + assert section.startswith(system), section + assert section.count(_INSTRUCTION) == 1, section + + +def test_two_tools_are_both_listed_under_one_instruction(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + _post_chat(gateway, model, _first_turn(), tools=[_WEATHER_TOOL, _TIME_TOOL]) + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather", "get_time") + assert prompt.index("'name': 'get_weather'") < prompt.index("'name': 'get_time'"), prompt + + +def test_unauthenticated_tool_request_never_reaches_ollama(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + code, text = _post( + gateway, + "/v1/chat/completions", + {"model": model, "messages": _first_turn(), "tools": [_WEATHER_TOOL]}, + key=f"sk-not-a-key-{uuid.uuid4().hex}", + ) + assert code == 401, text + assert _generate_calls(wire) == () + + +def test_ollama_model_not_found_reaches_the_caller_after_one_attempt(gateway: Gateway) -> None: + message: Final = f"model '{_BACKEND}' not found {uuid.uuid4().hex}" + reply: Final = Reply(status=404, body=json.dumps({"error": message}).encode()) + with _ollama_server(lambda _: reply) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + code, text = _post( + gateway, "/v1/chat/completions", {"model": model, "messages": _first_turn(), "tools": [_WEATHER_TOOL]} + ) + assert code == 404, text + assert message in text, text + _assert_instructed_once(_prompt_of(_only_generate(wire)), "get_weather") + + +def test_ollama_server_error_on_the_tool_result_turn_does_not_take_the_deployment_down(gateway: Gateway) -> None: + attempts: Final = itertools.count() + failure: Final = f"internal failure {uuid.uuid4().hex}" + + def respond(_: Request) -> Reply: + if next(attempts) == 0: + return Reply(status=500, body=json.dumps({"error": failure}).encode()) + return _generate_reply(_ANSWER) + + with _ollama_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + code, text = _post( + gateway, "/v1/chat/completions", {"model": model, "messages": _second_turn(), "tools": [_WEATHER_TOOL]} + ) + assert code == 500, text + assert failure in text, text + payload: Final = _post_chat(gateway, model, _second_turn(), tools=[_WEATHER_TOOL]) + choices: Final = payload["choices"] + assert isinstance(choices, list) and len(choices) == 1 + assert json.dumps(choices[0]).count(_ANSWER) == 1, choices + prompts: Final = [_prompt_of(_JSON_OBJECT.validate_json(request.body)) for request in _generate_calls(wire)] + assert len(prompts) == 2, prompts + for prompt in prompts: + _assert_tool_turn(prompt) + + +@pytest.mark.parametrize( + ("result", "forwarded"), + [ + pytest.param("", None, id="empty-string-drops-the-section"), + pytest.param("r" * 5120, "r" * 5120, id="5kb-string-forwarded-intact"), + pytest.param( + [{"type": "text", "text": "Paris: 22 degrees"}, {"type": "text", "text": "clear skies"}], + "Paris: 22 degreesclear skies", + id="text-parts-joined", + ), + ], +) +def test_tool_result_content_shapes_reach_the_prompt( + gateway: Gateway, result: JsonValue, forwarded: str | None +) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + payload: Final = _post_chat(gateway, model, _second_turn(result), tools=[_WEATHER_TOOL]) + assert json.dumps(payload["choices"]).count(_ANSWER) == 1, payload + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + if forwarded is None: + assert f"### User:\n{_QUESTION}\n\n### Assistant:\n{_CALL_JSON}\n\n### System:\n" in prompt, prompt + else: + _assert_tool_turn(prompt, forwarded) + + +def test_the_same_tool_result_twice_is_forwarded_twice_under_one_instruction(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + messages: Final = [*_second_turn(), {"role": "tool", "tool_call_id": _CALL_ID, "content": _RESULT}] + _post_chat(gateway, model, messages, tools=[_WEATHER_TOOL]) + prompt: Final = _prompt_of(_only_generate(wire)) + _assert_instructed_once(prompt, "get_weather") + assert prompt.count(_RESULT) == 2, prompt + assert prompt.count("### User:") == 2 and prompt.count("### Assistant:") == 1, prompt + + +def test_a_second_function_call_after_the_result_is_surfaced_as_tool_calls(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + completion: Final = _openai_client(gateway).chat.completions.create( + model=model, + messages=_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_WEATHER_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) + choice: Final = completion.choices[0] + assert choice.finish_reason == "tool_calls" + assert choice.message.tool_calls is not None and [call.function.name for call in choice.message.tool_calls] == [ + "get_weather" + ] + _assert_tool_turn(_prompt_of(_only_generate(wire))) + assert _spend_row(completion.id) == _billed(model) + + +def test_a_non_function_json_answer_is_returned_as_text(gateway: Gateway) -> None: + answer: Final = json.dumps({"city": "Paris", "temperature_c": 22, "sky": "clear"}) + with _ollama_server(lambda _: _generate_reply(answer)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + completion: Final = _openai_client(gateway).chat.completions.create( + model=model, + messages=_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_WEATHER_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) + choice: Final = completion.choices[0] + assert choice.finish_reason == "stop" + assert choice.message.tool_calls is None + assert choice.message.content is not None and json.loads(choice.message.content) == json.loads(answer) + _assert_tool_turn(_prompt_of(_only_generate(wire))) + assert _spend_row(completion.id) == _billed(model) + + +def test_int_tool_result_content_fails_in_the_response_body_and_leaves_the_deployment_serving(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + code, text = _post( + gateway, "/v1/chat/completions", {"model": model, "messages": _second_turn(22), "tools": [_WEATHER_TOOL]} + ) + assert code >= 400, text + error: Final = _JSON_OBJECT.validate_json(text)["error"] + assert isinstance(error, dict) and isinstance(error["message"], str) and error["message"], text + assert _generate_calls(wire) == () + payload: Final = _post_chat(gateway, model, _second_turn(), tools=[_WEATHER_TOOL]) + assert json.dumps(payload["choices"]).count(_ANSWER) == 1, payload + _assert_tool_turn(_prompt_of(_only_generate(wire))) + + +def test_ollama_chat_keeps_native_tools_and_gets_no_instruction(gateway: Gateway) -> None: + reply: Final = Reply( + body=json.dumps( + { + "model": _BACKEND, + "created_at": "2026-10-07T00:00:00Z", + "message": {"role": "assistant", "content": _ANSWER}, + "done": True, + "prompt_eval_count": 30, + "eval_count": 12, + } + ).encode() + ) + with wire_server(lambda _: reply) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama_chat/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + completion: Final = _openai_client(gateway).chat.completions.create( + model=model, + messages=_second_turn(), # pyright: ignore[reportArgumentType] # plain JSON messages + tools=[_WEATHER_TOOL], # pyright: ignore[reportArgumentType] # plain JSON tool + extra_body={"cache": _NO_CACHE}, + ) + assert completion.choices[0].message.content == _ANSWER + received: Final = tuple(request for request in wire.drain() if request.method == "POST") + assert [request.target for request in received] == ["/api/chat"] + body: Final = _JSON_OBJECT.validate_json(received[0].body) + assert "prompt" not in body and "format" not in body, sorted(body) + tools: Final = body["tools"] + assert isinstance(tools, list) and len(tools) == 1 + messages: Final = body["messages"] + assert isinstance(messages, list) and [item["role"] for item in messages if isinstance(item, dict)] == [ + "user", + "assistant", + "tool", + ] + assert "function_name" not in received[0].body.decode(), received[0].body + assert _spend_row(completion.id) == _billed(model) + + +def test_identical_uncached_tool_result_turns_are_each_forwarded_and_billed_once(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_ANSWER)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + first: Final = _post_chat(gateway, model, _second_turn(), tools=[_WEATHER_TOOL]) + second: Final = _post_chat(gateway, model, _second_turn(), tools=[_WEATHER_TOOL]) + assert first["id"] != second["id"], (first["id"], second["id"]) + prompts: Final = [_prompt_of(_JSON_OBJECT.validate_json(request.body)) for request in _generate_calls(wire)] + assert len(prompts) == 2, prompts + for prompt in prompts: + _assert_tool_turn(prompt) + for payload in (first, second): + identity: Final = payload["id"] + assert isinstance(identity, str) + assert _spend_row(identity) == _billed(model) + + +def test_the_cell_deployment_is_gone_after_its_scenario(gateway: Gateway) -> None: + with _ollama_server(lambda _: _generate_reply(_CALL_JSON)) as wire: + with gateway.scenario() as scenario: + model: Final = scenario.model(model=f"ollama/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + _post_chat(gateway, model, _first_turn(), tools=[_WEATHER_TOOL]) + _only_generate(wire) + listed: Final = eventually( + lambda: [entry["model_name"] for entry in _deployments(gateway) if isinstance(entry, dict)], + lambda names: model not in names, + seconds=70, + ) + assert model not in listed + + +def _deployments(gateway: Gateway) -> list[JsonValue]: + data: Final = gateway.get("/model/info")["data"] + assert isinstance(data, list) + return data diff --git a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py index a2987730851..61b59f7dbe6 100644 --- a/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py +++ b/tests/unit/litellm_core_utils/prompt_templates/test_litellm_core_utils_prompt_templates_factory.py @@ -51,6 +51,17 @@ def test_function_call_prompt_preserves_append_failure_for_non_string_content() function_call_prompt(messages, []) +def test_function_call_prompt_lets_the_model_answer_after_a_function_result() -> None: + messages: Final[list[dict[str, object]]] = [{"role": "system", "content": "Be terse."}] + + prompted: Final = function_call_prompt(messages, [{"name": "get_weather"}]) + + system: Final = str(prompted[0]["content"]) + assert "JSON OUTPUT ONLY" not in system + assert "reply to the user in plain text instead of calling a function again" in system + assert "{'name': 'get_weather'}" in system + + @pytest.mark.parametrize( ("thought_signature", "expected"), [ diff --git a/tests/unit/llms/base_llm/chat/test_attribution_headers.py b/tests/unit/llms/base_llm/chat/test_attribution_headers.py new file mode 100644 index 00000000000..f31b4e846bf --- /dev/null +++ b/tests/unit/llms/base_llm/chat/test_attribution_headers.py @@ -0,0 +1,210 @@ +""" +Provider attribution headers (`BaseConfig.get_attribution_headers`) must reach +the outbound request on every OpenAI-compatible chat path, and a caller header +with the same name must win. + +Requests go through a real `litellm.completion` into an in-process httpx +transport that records what would have been sent, because the default path +(OpenAI SDK) never calls `validate_environment`. +""" + +import json +from collections.abc import AsyncIterable, Iterable +from typing import Final, cast + +import httpx +import openai +import pytest + +import litellm +from litellm.llms.base_llm.chat.transformation import with_attribution_headers +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, HTTPHandler + +_API_BASE: Final = "https://provider.invalid/v1" +_COMPLETION_BODY: Final = json.dumps( + { + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 1, + "model": "m", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}, + } +).encode() +_STREAM_BODY: Final = ( + "data: " + + json.dumps( + { + "id": "chatcmpl-1", + "object": "chat.completion.chunk", + "created": 1, + "model": "m", + "choices": [{"index": 0, "delta": {"role": "assistant", "content": "hi"}, "finish_reason": "stop"}], + } + ) + + "\n\ndata: [DONE]\n\n" +).encode() + + +class _HeaderCapturingTransport(httpx.BaseTransport, httpx.AsyncBaseTransport): + """Records each outbound request's headers and answers like a chat completions server.""" + + def __init__(self) -> None: + self.sent: tuple[httpx.Headers, ...] = () + + def handle_request(self, request: httpx.Request) -> httpx.Response: + return self._respond(request, request.read()) + + async def handle_async_request(self, request: httpx.Request) -> httpx.Response: + return self._respond(request, await request.aread()) + + def _respond(self, request: httpx.Request, body: bytes) -> httpx.Response: + self.sent = (*self.sent, request.headers) + if json.loads(body).get("stream"): + return httpx.Response(200, content=_STREAM_BODY, headers={"content-type": "text/event-stream"}) + return httpx.Response(200, content=_COMPLETION_BODY, headers={"content-type": "application/json"}) + + def last(self, header: str) -> list[str]: + return self.sent[-1].get_list(header) + + +def _client(transport: _HeaderCapturingTransport, path: str, is_async: bool) -> object: + if path == "sdk": + if is_async: + return openai.AsyncOpenAI( + api_key="k", base_url=_API_BASE, http_client=httpx.AsyncClient(transport=transport) + ) + return openai.OpenAI(api_key="k", base_url=_API_BASE, http_client=httpx.Client(transport=transport)) + if is_async: + return AsyncHTTPHandler(transport=transport) + return HTTPHandler(client=httpx.Client(transport=transport)) + + +@pytest.fixture +def transport() -> _HeaderCapturingTransport: + return _HeaderCapturingTransport() + + +@pytest.fixture(params=["sdk", "http_handler"]) +def handler_path(request: pytest.FixtureRequest, monkeypatch: pytest.MonkeyPatch) -> str: + if request.param == "http_handler": + monkeypatch.setenv("EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER", "true") + else: + monkeypatch.delenv("EXPERIMENTAL_OPENAI_BASE_LLM_HTTP_HANDLER", raising=False) + return request.param + + +_NOVITA_MODEL: Final = "novita/meta-llama/llama-3.3-70b-instruct" + +_ATTRIBUTED: Final = [ + pytest.param(_NOVITA_MODEL, "x-novita-source", id="novita"), + pytest.param("perplexity/sonar", "x-pplx-integration", id="perplexity"), +] + + +def _drain(response: object) -> None: + for _ in cast(Iterable[object], response): + pass + + +async def _adrain(response: object) -> None: + async for _ in cast(AsyncIterable[object], response): + pass + + +@pytest.mark.parametrize(("model", "header"), _ATTRIBUTED) +@pytest.mark.parametrize("stream", [False, True]) +def test_attribution_header_sent_sync( + transport: _HeaderCapturingTransport, handler_path: str, model: str, header: str, stream: bool +) -> None: + response: Final = litellm.completion( + model=model, + messages=[{"role": "user", "content": "hi"}], + api_base=_API_BASE, + api_key="k", + stream=stream, + client=_client(transport, handler_path, is_async=False), + ) + if stream: + _drain(response) + + assert transport.last(header) == ["litellm"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize(("model", "header"), _ATTRIBUTED) +@pytest.mark.parametrize("stream", [False, True]) +async def test_attribution_header_sent_async( + transport: _HeaderCapturingTransport, handler_path: str, model: str, header: str, stream: bool +) -> None: + response: Final = await litellm.acompletion( + model=model, + messages=[{"role": "user", "content": "hi"}], + api_base=_API_BASE, + api_key="k", + stream=stream, + client=_client(transport, handler_path, is_async=True), + ) + if stream: + await _adrain(response) + + assert transport.last(header) == ["litellm"] + + +@pytest.mark.parametrize(("model", "header"), _ATTRIBUTED) +@pytest.mark.parametrize("header_kwarg", ["headers", "extra_headers"]) +def test_caller_header_overrides_attribution_any_casing( + transport: _HeaderCapturingTransport, handler_path: str, model: str, header: str, header_kwarg: str +) -> None: + caller_headers: Final = {header.upper(): "my-app"} + + litellm.completion( + model=model, + messages=[{"role": "user", "content": "hi"}], + api_base=_API_BASE, + api_key="k", + client=_client(transport, handler_path, is_async=False), + **{header_kwarg: caller_headers}, + ) + + assert transport.last(header) == ["my-app"] + assert caller_headers == {header.upper(): "my-app"} + + +def test_provider_without_attribution_sends_none(transport: _HeaderCapturingTransport, handler_path: str) -> None: + litellm.completion( + model="deepinfra/meta-llama/Meta-Llama-3-8B-Instruct", + messages=[{"role": "user", "content": "hi"}], + api_base=_API_BASE, + api_key="k", + client=_client(transport, handler_path, is_async=False), + ) + + assert transport.last("x-novita-source") == [] + assert transport.last("x-pplx-integration") == [] + + +def test_global_litellm_headers_still_apply_and_are_not_mutated( + transport: _HeaderCapturingTransport, handler_path: str, monkeypatch: pytest.MonkeyPatch +) -> None: + global_headers: Final = {"X-Global": "1"} + monkeypatch.setattr(litellm, "headers", global_headers) + + litellm.completion( + model=_NOVITA_MODEL, + messages=[{"role": "user", "content": "hi"}], + api_base=_API_BASE, + api_key="k", + client=_client(transport, handler_path, is_async=False), + ) + + assert transport.last("x-global") == ["1"] + assert transport.last("x-novita-source") == ["litellm"] + assert global_headers == {"X-Global": "1"} + + +def test_with_attribution_headers_returns_headers_unchanged_when_nothing_to_add() -> None: + headers: Final = {"A": "1"} + + assert with_attribution_headers({}, headers) is headers + assert with_attribution_headers({}, None) is None diff --git a/tests/unit/proxy/lens/test_endpoints.py b/tests/unit/proxy/lens/test_endpoints.py index c68b1449dc8..7d85c8e22ec 100644 --- a/tests/unit/proxy/lens/test_endpoints.py +++ b/tests/unit/proxy/lens/test_endpoints.py @@ -1,4 +1,5 @@ -from collections.abc import Callable +import asyncio +from collections.abc import Callable, Mapping from datetime import datetime, timedelta, timezone from types import SimpleNamespace from typing import Final @@ -6,21 +7,25 @@ from typing import Final import httpx import pytest from fastapi import HTTPException -from pydantic import ValidationError +from pydantic import TypeAdapter, ValidationError import litellm from litellm import Router from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.lens.endpoints import ( claim_due, + get_signals, list_agents, + put_signals, read_reviews, result, run_settings, run_window, trace_findings, + trace_signal_statuses, user_scope, validate_model, + validate_signal_model, watchable, watching, worker_supports_model, @@ -45,6 +50,7 @@ from litellm.proxy.lens.models import ( Worker, ) from litellm.proxy.lens.repository import DueLens, Row +from litellm.proxy.lens.signals import SignalConfig, StoredTraceSignal from litellm.proxy.lens.state import claim_job, queue_job, replace_job from litellm.rust_bridge.trace.generated.models import ExecutionRow, LensSampleParams from litellm.rust_bridge.trace.storage import ClickHouseStorage @@ -80,6 +86,46 @@ class ResultDatabase: return len(self.completed) +class SignalStatusDatabase: + def __init__(self, config: SignalConfig, rows: Mapping[str, StoredTraceSignal]) -> None: + self.config: Final = config + self.rows: Final = rows + self.saved: Final[asyncio.Queue[tuple[object, ...]]] = asyncio.Queue() + + async def query_raw(self, query: str, *args: object) -> object: + if '"LiteLLM_LensSignalConfig"' in query: + return ({"data": self.config.model_dump(mode="json")},) + payload: Final = args[0] + assert isinstance(payload, str) + requested: Final = TypeAdapter(tuple[TraceIdentity, ...]).validate_json(payload) + return tuple( + {"data": row.model_dump(mode="json")} + for identity in requested + if (row := self.rows.get(identity.trace_id)) is not None + ) + + async def execute_raw(self, query: str, *args: object) -> int: + await self.saved.put(args) + return 1 + + +def signal_router() -> Router: + return Router( + model_list=[ + { + "model_name": "decision", + "litellm_params": {"model": "openai/test-decision", "api_key": "test-key"}, + "model_info": {"mode": "evaluation"}, + }, + { + "model_name": "chat", + "litellm_params": {"model": "openai/test-chat", "api_key": "test-key"}, + "model_info": {"mode": "chat"}, + }, + ] + ) + + @pytest.mark.asyncio @pytest.mark.parametrize("change", ("cancelled", "expired", "reclaimed", "reassigned")) async def test_result_cannot_commit_after_losing_ownership_during_evidence_validation( @@ -526,6 +572,144 @@ async def test_trace_finding_counts_require_investigation_read_access() -> None: assert error.value.status_code == 403 +@pytest.mark.asyncio +async def test_signal_endpoints_return_statuses_in_request_order_for_admin_viewers( + monkeypatch: pytest.MonkeyPatch, +) -> None: + from litellm.proxy import proxy_server + + config: Final = SignalConfig(model="decision") + rows: Final = { + "pending": StoredTraceSignal( + trace_id="pending", + config_key=config.key(), + span_count=1, + claimed_until=NOW + timedelta(minutes=1), + data={"status": "pending", "scores": {}, "model": "decision", "error": ""}, + ), + "classified": StoredTraceSignal( + trace_id="classified", + config_key=config.key(), + span_count=1, + classified_at=NOW, + data={ + "status": "classified", + "scores": {"user_frustration": 0.7, "missing_capability": 0.8}, + "model": "decision", + "error": "", + }, + ), + "failed": StoredTraceSignal( + trace_id="failed", + config_key=config.key(), + span_count=1, + classified_at=NOW, + data={"status": "failed", "scores": {}, "model": "decision", "error": "classification failed"}, + ), + "stale": StoredTraceSignal( + trace_id="stale", + config_key="old-config", + span_count=1, + classified_at=NOW, + data={ + "status": "classified", + "scores": {"user_frustration": 1.0}, + "model": "old", + "error": "", + }, + ), + } + database: Final = SignalStatusDatabase(config, rows) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=database)) + viewer: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) + request: Final = TraceFindingsRequest( + traces=tuple( + TraceIdentity(trace_id=trace_id) for trace_id in ("failed", "classified", "missing", "pending", "stale") + ) + ) + + assert await get_signals(viewer) == config + results: Final = await trace_signal_statuses(request, viewer) + + assert tuple((result.trace_id, result.status) for result in results) == ( + ("failed", "failed"), + ("classified", "classified"), + ("missing", "unclassified"), + ("pending", "pending"), + ("stale", "unclassified"), + ) + assert tuple((flag.signal_id, flag.name, flag.score) for flag in results[1].flags) == ( + ("missing_capability", "Missing capability", 0.8), + ("user_frustration", "User frustration", 0.7), + ) + + +@pytest.mark.asyncio +async def test_signal_endpoints_require_connected_postgres(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "prisma_client", None) + auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) + + with pytest.raises(HTTPException) as error: + await get_signals(auth) + + assert error.value.status_code == 503 + + +@pytest.mark.asyncio +async def test_put_signals_saves_config_for_admin(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + + database: Final = SignalStatusDatabase(SignalConfig(), {}) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=database)) + monkeypatch.setattr(proxy_server, "llm_router", signal_router()) + auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + body: Final = SignalConfig(model="decision", threshold=0.7) + + assert await put_signals(body, auth) == body + + saved: Final = await database.saved.get() + assert saved[0] == "global" + assert isinstance(saved[1], str) + assert SignalConfig.model_validate_json(saved[1]) == body + + +def test_signal_model_requires_a_ready_router() -> None: + with pytest.raises(HTTPException) as error: + validate_signal_model(SignalConfig(model="decision"), None) + + assert error.value.status_code == 400 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("role", (LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY)) +async def test_put_signals_rejects_non_admin_roles(role: LitellmUserRoles) -> None: + auth: Final = UserAPIKeyAuth(user_role=role) + with pytest.raises(HTTPException) as error: + await put_signals(SignalConfig(model="decision"), auth) + assert error.value.status_code == 403 + + +@pytest.mark.asyncio +@pytest.mark.parametrize("model", ("chat", "unconfigured")) +async def test_put_signals_rejects_chat_and_unknown_model_groups(monkeypatch: pytest.MonkeyPatch, model: str) -> None: + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "llm_router", signal_router()) + auth: Final = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) + + with pytest.raises(HTTPException) as error: + await put_signals(SignalConfig(model=model), auth) + + assert error.value.status_code == 400 + assert error.value.detail == "Choose a System 1 model (evaluation mode) configured on this proxy" + + +def test_signal_model_accepts_only_evaluation_mode_groups() -> None: + assert validate_signal_model(SignalConfig(model="decision"), signal_router()) is None + + @pytest.mark.parametrize( "role", (LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, LitellmUserRoles.TEAM), diff --git a/tests/unit/proxy/lens/test_signals.py b/tests/unit/proxy/lens/test_signals.py new file mode 100644 index 00000000000..d2c4eda3966 --- /dev/null +++ b/tests/unit/proxy/lens/test_signals.py @@ -0,0 +1,1002 @@ +import asyncio +import json +from collections.abc import AsyncGenerator, Mapping, Sequence +from contextlib import asynccontextmanager +from datetime import datetime, timedelta, timezone +from itertools import chain +from types import MappingProxyType, SimpleNamespace +from typing import Final + +import pytest +from pydantic import JsonValue, TypeAdapter, ValidationError + +from litellm.proxy.lens.models import Execution, Scope, TraceIdentity +from litellm.proxy.lens.repository import Database, Row +from litellm.proxy.lens.signal_repository import SignalRepository +from litellm.proxy.lens.signals import ( + DEFAULT_SIGNALS, + SIGNAL_CLAIM_LEASE, + SIGNAL_MAX_SCAN_PAGES, + SIGNAL_TASK, + DecisionQuestions, + DecisionState, + Signal, + SignalAttempt, + SignalClassifier, + SignalConfig, + SignalData, + SignalStep, + StoredTraceSignal, + candidate, + run_signal_loop, + run_signal_tick, + signal_state, + trace_signals, +) +from litellm.proxy.lens.sources import SourceReader +from litellm.rust_bridge.trace.generated.models import ( + ActivityAvailability, + AgentRow, + CountRow, + ExecutionRow, + LensAccessParams, + LensContentParams, + LensEvidenceParams, + LensSampleParams, + PartRow, +) +from litellm.types.decisions import DecisionsResponse +from litellm.types.decisions import NoulAnswer as DecisionsNoulAnswer + +NOW: Final = datetime(2026, 10, 7, 12, tzinfo=timezone.utc) +CURRENT_CONFIG_KEY: Final = SignalConfig(model="decision").key() +_SIGNAL_STEPS: Final[TypeAdapter[tuple[SignalStep, ...]]] = TypeAdapter(tuple[SignalStep, ...]) +_STORED_DATA: Final[TypeAdapter[JsonValue]] = TypeAdapter(JsonValue) + + +def execution(identity: str, span_count: int = 1) -> Execution: + return Execution( + id=identity, + source="traces", + trace_id=identity, + team_id="", + name=identity, + start_time="", + span_count=span_count, + ) + + +def part(identity: str, content: str) -> PartRow: + return PartRow( + span_id=identity, + parent_span_id="", + name=identity, + kind="agent", + start_time="", + end_time="", + content=content, + truncated=0, + ) + + +def stored_trace( + config_key: str, + *, + trace_id: str = "trace", + status: str = "classified", + span_count: int = 1, + claimed_until: datetime | None = None, + classified_at: datetime | None = NOW - timedelta(minutes=10), + scores: dict[str, float] | None = None, + error: str = "", +) -> StoredTraceSignal: + return StoredTraceSignal( + trace_id=trace_id, + trace_ref="", + config_key=config_key, + span_count=span_count, + claimed_until=claimed_until, + classified_at=classified_at, + data=_STORED_DATA.validate_python( + { + "status": status, + "scores": scores or {}, + "model": "decision", + "error": error, + } + ), + ) + + +class SignalStorage: + def __init__( + self, + executions: tuple[ExecutionRow, ...] = (), + parts: tuple[PartRow, ...] = (), + ) -> None: + self.executions: Final = executions + self.parts: Final = parts + + async def lens_availability(self, parameters: LensAccessParams) -> Sequence[ActivityAvailability]: + return () + + async def lens_agents(self, parameters: LensAccessParams) -> Sequence[AgentRow]: + return () + + async def lens_sample(self, parameters: LensSampleParams) -> Sequence[ExecutionRow]: + return self.executions + + async def lens_content(self, parameters: LensContentParams) -> Sequence[PartRow]: + return self.parts or (part(parameters.id, parameters.id),) + + async def lens_evidence(self, parameters: LensEvidenceParams) -> Sequence[CountRow]: + return () + + +class PagedSignalStorage(SignalStorage): + def __init__(self, pages: tuple[tuple[PartRow, ...], ...]) -> None: + super().__init__() + self.pages: Final = pages + + async def lens_content(self, parameters: LensContentParams) -> Sequence[PartRow]: + index: Final = int(parameters.cursor) if parameters.cursor else 0 + return self.pages[index] + + +class PagedSampleStorage(SignalStorage): + def __init__(self, pages: tuple[tuple[ExecutionRow, ...], ...], initial_cursor: str = "") -> None: + super().__init__() + self.pages: Final = pages + self.cursors: Final[asyncio.Queue[str]] = asyncio.Queue() + self.page_by_cursor: Final = MappingProxyType( + { + initial_cursor: 0, + **{page[-1].selection_key: index + 1 for index, page in enumerate(pages[:-1])}, + } + ) + + async def lens_sample(self, parameters: LensSampleParams) -> Sequence[ExecutionRow]: + await self.cursors.put(parameters.after) + index: Final = self.page_by_cursor[parameters.after] + return self.pages[index] + + +class SignalDatabase: + def __init__( + self, + config: SignalConfig | None, + *, + stored_rows: tuple[StoredTraceSignal, ...] = (), + claim_result: bool = True, + ) -> None: + self.config: Final = config + self.stored_rows: Final = stored_rows + self.claim_result: Final = claim_result + self.calls: Final[asyncio.Queue[str]] = asyncio.Queue() + self.claims: Final[asyncio.Queue[str]] = asyncio.Queue() + self.claim_args: Final[asyncio.Queue[tuple[object, ...]]] = asyncio.Queue() + self.saved: Final[asyncio.Queue[tuple[object, ...]]] = asyncio.Queue() + + async def query_raw(self, query: str, *args: object) -> object: + if '"LiteLLM_LensSignalConfig"' in query: + return () if self.config is None else (Row(data=self.config.model_dump(mode="json")),) + if query.startswith("SELECT jsonb_build_object"): + payload: Final = args[0] + assert isinstance(payload, str) + requested: Final = TypeAdapter(tuple[TraceIdentity, ...]).validate_json(payload) + identities: Final = tuple((trace.trace_id, trace.trace_ref) for trace in requested) + return tuple( + Row(data=stored.model_dump(mode="json")) + for stored in self.stored_rows + if (stored.trace_id, stored.trace_ref) in identities + ) + if query.startswith('INSERT INTO "LiteLLM_LensTraceSignal"'): + await self.claim_args.put(args) + if not self.claim_result: + return () + trace_id: Final = args[0] + assert isinstance(trace_id, str) + await self.claims.put(trace_id) + return (Row(data={"trace_id": trace_id}),) + raise AssertionError(f"Unexpected query: {query}") + + async def execute_raw(self, query: str, *args: object) -> int: + await self.saved.put(args) + return 1 + + @asynccontextmanager + async def transaction(self) -> AsyncGenerator[Database, None]: + yield self + + +def saved_result(args: tuple[object, ...]) -> SignalData: + payload: Final = args[1] + assert isinstance(payload, str) + return SignalData.model_validate_json(payload) + + +@pytest.mark.asyncio +async def test_signal_repository_reads_defaults_and_saves_the_global_config() -> None: + database: Final = SignalDatabase(None) + repository: Final = SignalRepository(database) + updated: Final = SignalConfig(model="decision", threshold=0.7) + + assert await repository.get_config() == SignalConfig() + await repository.save_config(updated) + + saved: Final = await database.saved.get() + assert saved[0] == "global" + assert isinstance(saved[1], str) + assert SignalConfig.model_validate_json(saved[1]) == updated + + +@pytest.mark.asyncio +async def test_signal_repository_reads_rows_and_reports_a_lost_claim() -> None: + config: Final = SignalConfig(model="decision") + row: Final = stored_trace(config.key()) + database: Final = SignalDatabase(config, stored_rows=(row,), claim_result=False) + repository: Final = SignalRepository(database) + + assert await repository.traces(()) == () + assert await repository.traces((TraceIdentity(trace_id="trace"),)) == (row,) + assert not await repository.claim(execution("trace"), config, NOW + timedelta(minutes=5), NOW) + + +def test_signal_config_hashes_questions_but_not_threshold_or_display_name() -> None: + config: Final = SignalConfig(model="decision") + different_threshold: Final = config.model_copy(update={"threshold": 0.9}) + renamed: Final = config.model_copy( + update={ + "signals": ( + config.signals[0].model_copy(update={"name": "Frustration"}), + *config.signals[1:], + ) + } + ) + changed_question: Final = config.model_copy( + update={ + "signals": ( + config.signals[0].model_copy(update={"question": "Does this user sound upset?"}), + *config.signals[1:], + ) + } + ) + + assert config.key() == different_threshold.key() == renamed.key() + assert config.key() != changed_question.key() + assert DEFAULT_SIGNALS == config.signals + + +def test_signal_config_rejects_duplicate_ids_and_non_finite_thresholds() -> None: + duplicate: Final = Signal(id="same", name="First", question="Question one") + with pytest.raises(ValidationError): + SignalConfig(signals=(duplicate, duplicate)) + with pytest.raises(ValidationError): + SignalConfig(threshold=float("nan")) + + +@pytest.mark.parametrize( + "stored,trace_count,expected", + ( + (None, 1, True), + (stored_trace("old"), 1, True), + ( + stored_trace(CURRENT_CONFIG_KEY, span_count=1, classified_at=NOW - timedelta(minutes=6)), + 2, + True, + ), + ( + stored_trace( + CURRENT_CONFIG_KEY, + status="failed", + classified_at=(NOW - timedelta(minutes=31)).replace(tzinfo=None), + ), + 1, + True, + ), + (stored_trace(CURRENT_CONFIG_KEY), 1, False), + ( + stored_trace(CURRENT_CONFIG_KEY, span_count=1, classified_at=NOW - timedelta(minutes=2)), + 2, + False, + ), + ( + stored_trace("old", claimed_until=NOW + timedelta(minutes=1)), + 1, + False, + ), + ( + stored_trace( + CURRENT_CONFIG_KEY, + status="pending", + claimed_until=(NOW + timedelta(minutes=1)).replace(tzinfo=None), + classified_at=None, + ), + 1, + False, + ), + ( + stored_trace( + CURRENT_CONFIG_KEY, + status="pending", + claimed_until=(NOW - timedelta(minutes=1)).replace(tzinfo=None), + classified_at=None, + ), + 1, + True, + ), + (stored_trace(CURRENT_CONFIG_KEY, span_count=2), 1, False), + (stored_trace(CURRENT_CONFIG_KEY, span_count=1, classified_at=None), 2, False), + ), +) +def test_candidate_selection_respects_config_span_age_failure_age_and_claims( + stored: StoredTraceSignal | None, trace_count: int, expected: bool +) -> None: + config: Final = SignalConfig(model="decision") + assert candidate(execution("trace", trace_count), stored, config.key(), NOW) is expected + + +@pytest.mark.asyncio +async def test_classifier_sends_noul_questions_and_keeps_every_signal_score() -> None: + config: Final = SignalConfig(model="decision") + run: Final = execution("trace") + storage: Final = SignalStorage(parts=(part("agent", "user asks for a result"),)) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + assert model == "decision" + assert state == { + "task": SIGNAL_TASK, + "steps": ({"kind": "agent", "name": "agent", "content": "user asks for a result"},), + } + assert _STORED_DATA.validate_json(json.dumps(state)) == { + "task": SIGNAL_TASK, + "steps": [{"kind": "agent", "name": "agent", "content": "user asks for a result"}], + } + expected_questions: Final = { + signal.id: {"type": "noul", "instructions": signal.question} for signal in config.signals + } + assert questions == expected_questions + assert _STORED_DATA.validate_json(json.dumps(questions)) == expected_questions + assert timeout == 60 + assert metadata == {"tags": ["litellm-lens-signals"]} + return DecisionsResponse( + answers={ + "user_frustration": DecisionsNoulAnswer(type="noul", noul=0.9), + "missing_capability": DecisionsNoulAnswer(type="noul", noul=0.6), + "repeated_request": DecisionsNoulAnswer(type="noul", noul=0.2), + "unknown": DecisionsNoulAnswer(type="noul", noul=1.0), + } + ) + + attempt: Final = await SignalClassifier(SourceReader(storage), decide, lambda: NOW).classify( + Scope(all_teams=True), run, config + ) + + assert attempt == SignalAttempt( + status="classified", + scores={"user_frustration": 0.9, "missing_capability": 0.6, "repeated_request": 0.2}, + model="decision", + ) + + +@pytest.mark.asyncio +async def test_missing_noul_answer_fails_while_unknown_and_non_noul_answers_are_ignored() -> None: + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return { + "answers": { + "user_frustration": {"type": "noul", "noul": 0.9}, + "missing_capability": {"type": "choice", "choice": "yes"}, + "unknown": {"type": "noul", "noul": 1.0}, + } + } + + attempt: Final = await SignalClassifier( + SourceReader(SignalStorage(parts=(part("agent", "content"),))), + decide, + lambda: NOW, + ).classify(Scope(all_teams=True), execution("trace"), SignalConfig(model="decision")) + + assert attempt.status == "failed" + assert attempt.scores == {"user_frustration": 0.9} + assert attempt.error == "Decisions response omitted a configured noul answer" + + +@pytest.mark.asyncio +async def test_classifier_turns_decisions_errors_into_failed_attempts() -> None: + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + raise RuntimeError("decisions unavailable") + + attempt: Final = await SignalClassifier( + SourceReader(SignalStorage(parts=(part("agent", "content"),))), + decide, + lambda: NOW, + ).classify(Scope(all_teams=True), execution("trace"), SignalConfig(model="decision")) + + assert attempt == SignalAttempt(status="failed", model="decision", error="decisions unavailable") + + +def test_signal_flags_use_current_threshold_and_current_display_name() -> None: + config: Final = SignalConfig(model="decision") + row: Final = stored_trace( + config.key(), + scores={"user_frustration": 0.91, "missing_capability": 0.67, "repeated_request": 0.49}, + ) + high_threshold: Final = config.model_copy( + update={ + "threshold": 0.9, + "signals": ( + config.signals[0].model_copy(update={"name": "Frustrated user"}), + *config.signals[1:], + ), + } + ) + trace: Final = TraceIdentity(trace_id="trace") + lower: Final = trace_signals(trace, row, config) + higher: Final = trace_signals(trace, row, high_threshold) + + assert config.key() == high_threshold.key() + assert tuple((flag.signal_id, flag.score) for flag in lower.flags) == ( + ("user_frustration", 0.91), + ("missing_capability", 0.67), + ) + assert tuple((flag.signal_id, flag.name, flag.score) for flag in higher.flags) == ( + ("user_frustration", "Frustrated user", 0.91), + ) + assert not candidate(execution("trace"), row, high_threshold.key(), NOW) + + +def test_signal_flags_report_stored_errors_even_when_the_status_is_classified() -> None: + config: Final = SignalConfig(model="decision") + row: Final = stored_trace(config.key(), status="classified", error="classification failed") + + result: Final = trace_signals(TraceIdentity(trace_id="trace"), row, config) + + assert result.status == "failed" + assert result.model == "decision" + assert result.classified_at == row.classified_at + + +@pytest.mark.asyncio +async def test_signal_state_caps_content_to_head_and_tail_with_omitted_step() -> None: + parts: Final = tuple(part(str(index), chr(97 + index) * 2000) for index in range(30)) + state: Final = await signal_state( + SourceReader(SignalStorage(parts=parts)), + Scope(all_teams=True), + execution("trace"), + ) + steps_value: Final = state["steps"] + assert isinstance(steps_value, tuple) + steps: Final = _SIGNAL_STEPS.validate_python(steps_value) + head: Final = steps[:8] + marker: Final = steps[8] + tail: Final = steps[9:] + + assert state["task"] == SIGNAL_TASK + assert sum(len(step.content) for step in head) == 15000 + assert sum(len(step.content) for step in tail) == 25000 + assert head[0].content == "a" * 2000 + assert head[-1].content == "h" * 1000 + assert marker == SignalStep(kind="omitted", name="", content="9 steps omitted") + assert tail[0].content == "r" * 1000 + assert tail[-1].content == "~" * 2000 + + +@pytest.mark.asyncio +async def test_signal_state_limits_content_pages_and_part_sizes() -> None: + pages: Final = tuple( + tuple(part(f"page-{page}-{index}", "x" * 2501 if index == 0 else "x") for index in range(39)) + + (part(str(page + 1), "x"),) + for page in range(4) + ) + state: Final = await signal_state( + SourceReader(PagedSignalStorage(pages)), + Scope(all_teams=True), + execution("trace"), + ) + steps: Final = _SIGNAL_STEPS.validate_python(state["steps"]) + + assert len(steps) == 120 + assert steps[0].content.startswith("x" * 800) + assert "[... 501 characters omitted ...]" in steps[0].content + assert steps[0].content.endswith("x" * 1200) + assert steps[-1].name == "3" + assert all(not step.name.startswith("page-3-") for step in steps) + + small_state: Final = await signal_state( + SourceReader(SignalStorage(parts=(part("small", "ok"),))), + Scope(all_teams=True), + execution("trace"), + ) + small_steps: Final = _SIGNAL_STEPS.validate_python(small_state["steps"]) + assert small_steps == (SignalStep(kind="agent", name="small", content="ok"),) + + +@pytest.mark.asyncio +async def test_signal_state_part_excerpt_preserves_the_output_tail() -> None: + content: Final = "I" * 5000 + "OUTPUT: refused" + state: Final = await signal_state( + SourceReader(SignalStorage(parts=(part("result", content),))), + Scope(all_teams=True), + execution("trace"), + ) + steps: Final = _SIGNAL_STEPS.validate_python(state["steps"]) + excerpt: Final = steps[0].content + marker: Final = "\n[... 3015 characters omitted ...]\n" + + assert marker in excerpt + assert excerpt.endswith("OUTPUT: refused") + assert len(excerpt) == 800 + len(marker) + 1200 + + +@pytest.mark.asyncio +async def test_signal_tick_classifies_at_most_50_traces_and_persists_scores() -> None: + config: Final = SignalConfig(model="decision") + executions: Final = tuple( + ExecutionRow( + source="traces", + trace_id=f"trace-{index}", + team_id="", + name=f"trace-{index}", + start_time="", + span_count=1, + root_seen=1, + eligible=60, + selected=60, + selection_key=f"cursor-{index}", + ) + for index in range(60) + ) + storage: Final = SignalStorage(executions=executions) + database: Final = SignalDatabase(config) + repository: Final = SignalRepository(database) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + steps: Final = TypeAdapter(tuple[SignalStep, ...]).validate_python(state["steps"]) + await database.calls.put(steps[0].name) + return { + "answers": { + "user_frustration": {"type": "noul", "noul": 0.9}, + "missing_capability": {"type": "noul", "noul": 0.6}, + "repeated_request": {"type": "noul", "noul": 0.2}, + } + } + + await run_signal_tick(storage, repository, decide, lambda: NOW) + classified: Final = tuple(database.saved.get_nowait() for _ in range(database.saved.qsize())) + traces: Final = tuple(database.calls.get_nowait() for _ in range(database.calls.qsize())) + saved_data: Final = tuple(saved_result(args) for args in classified) + + assert len(classified) == 50 + assert frozenset(traces) == frozenset(f"trace-{index}" for index in range(50)) + assert ( + saved_data + == ( + SignalData( + status="classified", + scores={ + "user_frustration": 0.9, + "missing_capability": 0.6, + "repeated_request": 0.2, + }, + model="decision", + error="", + ), + ) + * 50 + ) + + +@pytest.mark.asyncio +async def test_signal_tick_claims_with_worker_start_time_and_skips_lost_claims() -> None: + config: Final = SignalConfig(model="decision") + executions: Final = tuple( + ExecutionRow( + source="traces", + trace_id=f"trace-{index}", + team_id="", + name=f"trace-{index}", + start_time="", + span_count=1, + root_seen=1, + eligible=2, + selected=2, + selection_key=f"cursor-{index}", + ) + for index in range(2) + ) + database: Final = SignalDatabase(config, claim_result=False) + repository: Final = SignalRepository(database) + + class AdvancingClock: + def __init__(self) -> None: + self.values: Final = tuple(NOW + timedelta(minutes=index) for index in range(3)) + self.index: int = 0 + + def __call__(self) -> datetime: + value: Final = self.values[self.index] + self.index += 1 + return value + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return {"answers": {}} + + await run_signal_tick(SignalStorage(executions=executions), repository, decide, AdvancingClock()) + + claims: Final = tuple(database.claim_args.get_nowait() for _ in range(database.claim_args.qsize())) + + def claim_times(args: tuple[object, ...]) -> tuple[datetime, datetime]: + claimed_until: Final = args[4] + claimed_at: Final = args[6] + assert isinstance(claimed_until, datetime) + assert isinstance(claimed_at, datetime) + return claimed_until, claimed_at + + times: Final = tuple(claim_times(claim) for claim in claims) + assert database.calls.empty() + assert database.saved.empty() + assert all(claimed_until == claimed_at + SIGNAL_CLAIM_LEASE for claimed_until, claimed_at in times) + assert all(claimed_at != NOW for _, claimed_at in times) + + +@pytest.mark.asyncio +async def test_signal_tick_resumes_after_ten_pages_and_resets_after_a_short_page() -> None: + config: Final = SignalConfig(model="decision") + + def sample_page(page: int) -> tuple[ExecutionRow, ...]: + return tuple( + ExecutionRow( + source="traces", + trace_id=f"trace-{page}-{index}", + team_id="", + name=f"trace-{page}-{index}", + start_time="", + span_count=1, + root_seen=1, + eligible=2500, + selected=2500, + selection_key=f"page-{page}-{index}", + ) + for index in range(100) + ) + + pages: Final = tuple(sample_page(page) for page in range(25)) + all_rows: Final = tuple(chain.from_iterable(pages)) + stored_rows: Final = tuple(stored_trace(CURRENT_CONFIG_KEY, trace_id=row.trace_id) for row in all_rows) + storage: Final = PagedSampleStorage(pages) + database: Final = SignalDatabase(config, stored_rows=stored_rows) + repository: Final = SignalRepository(database) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return {"answers": {}} + + first_cursor: Final = await run_signal_tick(storage, repository, decide, lambda: NOW) + first_calls: Final = tuple(storage.cursors.get_nowait() for _ in range(storage.cursors.qsize())) + second_cursor: Final = await run_signal_tick( + storage, + repository, + decide, + lambda: NOW, + cursor=first_cursor, + ) + second_calls: Final = tuple(storage.cursors.get_nowait() for _ in range(storage.cursors.qsize())) + + assert len(first_calls) == SIGNAL_MAX_SCAN_PAGES + assert first_cursor + assert len(second_calls) == SIGNAL_MAX_SCAN_PAGES + assert second_calls[0] == first_cursor + assert second_cursor + + short_storage: Final = PagedSampleStorage((pages[0][:50],)) + short_database: Final = SignalDatabase(config, stored_rows=stored_rows[:50]) + short_cursor: Final = await run_signal_tick( + short_storage, + SignalRepository(short_database), + decide, + lambda: NOW, + ) + assert short_cursor == "" + + +@pytest.mark.asyncio +async def test_signal_tick_resumes_a_partially_consumed_page() -> None: + config: Final = SignalConfig(model="decision") + page: Final = tuple( + ExecutionRow( + source="traces", + trace_id=f"trace-{index}", + team_id="", + name=f"trace-{index}", + start_time="", + span_count=1, + root_seen=1, + eligible=100, + selected=100, + selection_key=f"cursor-{index}", + ) + for index in range(100) + ) + initial_rows: Final = tuple(stored_trace(CURRENT_CONFIG_KEY, trace_id=f"trace-{index}") for index in range(20)) + resume_cursor: Final = "resume-page" + storage: Final = PagedSampleStorage((page, ()), initial_cursor=resume_cursor) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return { + "answers": { + "user_frustration": {"type": "noul", "noul": 0.9}, + "missing_capability": {"type": "noul", "noul": 0.6}, + "repeated_request": {"type": "noul", "noul": 0.2}, + } + } + + first_database: Final = SignalDatabase(config, stored_rows=initial_rows) + first_cursor: Final = await run_signal_tick( + storage, + SignalRepository(first_database), + decide, + lambda: NOW, + cursor=resume_cursor, + ) + first_claims: Final = tuple(first_database.claims.get_nowait() for _ in range(first_database.claims.qsize())) + + classified_first_rows: Final = tuple( + stored_trace(CURRENT_CONFIG_KEY, trace_id=trace_id) for trace_id in first_claims + ) + second_database: Final = SignalDatabase(config, stored_rows=(*initial_rows, *classified_first_rows)) + second_cursor: Final = await run_signal_tick( + storage, + SignalRepository(second_database), + decide, + lambda: NOW, + cursor=first_cursor, + ) + second_claims: Final = tuple(second_database.claims.get_nowait() for _ in range(second_database.claims.qsize())) + sample_cursors: Final = tuple(storage.cursors.get_nowait() for _ in range(storage.cursors.qsize())) + expected_eligible: Final = frozenset(f"trace-{index}" for index in range(20, 100)) + + assert first_cursor == resume_cursor + assert second_cursor == "" + assert len(first_claims) == 50 + assert len(second_claims) == 30 + assert frozenset(first_claims).isdisjoint(second_claims) + assert frozenset(first_claims) | frozenset(second_claims) == expected_eligible + assert sample_cursors == (resume_cursor, resume_cursor, page[-1].selection_key) + + +@pytest.mark.asyncio +async def test_signal_tick_skips_claims_and_writes_when_router_is_not_ready() -> None: + config: Final = SignalConfig(model="decision") + storage: Final = SignalStorage( + executions=( + ExecutionRow( + source="traces", + trace_id="trace", + team_id="", + name="trace", + start_time="", + span_count=1, + root_seen=1, + eligible=1, + selected=1, + selection_key="cursor", + ), + ) + ) + database: Final = SignalDatabase(config) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return {"answers": {}} + + await run_signal_tick( + storage, + SignalRepository(database), + decide, + lambda: NOW, + router_ready=lambda: False, + ) + + assert database.claims.empty() + assert database.saved.empty() + + +@pytest.mark.asyncio +async def test_signal_tick_skips_missing_dependencies_and_disabled_configs() -> None: + storage: Final = SignalStorage() + + await run_signal_tick(storage, None, None, lambda: NOW) + + database: Final = SignalDatabase(SignalConfig()) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + raise AssertionError("disabled signal config should not call Decisions") + + await run_signal_tick(storage, SignalRepository(database), decide, lambda: NOW) + assert database.claims.empty() + assert database.saved.empty() + + +class FailingStoreDatabase(SignalDatabase): + async def execute_raw(self, query: str, *args: object) -> int: + raise RuntimeError("store unavailable") + + +@pytest.mark.asyncio +async def test_signal_tick_continues_when_storing_a_result_fails() -> None: + config: Final = SignalConfig(model="decision") + execution_row: Final = ExecutionRow( + source="traces", + trace_id="trace", + team_id="", + name="trace", + start_time="", + span_count=1, + root_seen=1, + eligible=1, + selected=1, + selection_key="cursor", + ) + database: Final = FailingStoreDatabase(config) + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return { + "answers": { + "user_frustration": {"type": "noul", "noul": 0.9}, + "missing_capability": {"type": "noul", "noul": 0.6}, + "repeated_request": {"type": "noul", "noul": 0.2}, + } + } + + await run_signal_tick( + SignalStorage(executions=(execution_row,)), + SignalRepository(database), + decide, + lambda: NOW, + ) + + assert await database.claims.get() == "trace" + assert database.saved.empty() + + +class FailingSignalRepository: + def __init__(self) -> None: + self.started: Final = asyncio.Event() + + async def get_config(self) -> SignalConfig: + self.started.set() + await asyncio.sleep(0) + raise RuntimeError("tick failed") + + +@pytest.mark.asyncio +async def test_signal_loop_continues_after_a_tick_error() -> None: + repository: Final = FailingSignalRepository() + + async def decide( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return {"answers": {}} + + task: Final = asyncio.create_task(run_signal_loop(SignalStorage(), repository, decide, lambda: NOW)) + await repository.started.wait() + await asyncio.sleep(0) + task.cancel() + with pytest.raises(asyncio.CancelledError): + await task + + +@pytest.mark.asyncio +async def test_proxy_signal_call_resolves_the_current_router(monkeypatch: pytest.MonkeyPatch) -> None: + from litellm.proxy import proxy_server + + async def first_decisions( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return "first" + + async def second_decisions( + *, + model: str, + state: DecisionState, + questions: DecisionQuestions, + timeout: float, + metadata: Mapping[str, object], + ) -> object: + return "second" + + async def call_current_router() -> object: + return await proxy_server._call_current_lens_signal_router( + model="decision", + state={"task": "task"}, + questions={}, + timeout=60, + metadata={"tags": ["test"]}, + ) + + monkeypatch.setattr(proxy_server, "llm_router", SimpleNamespace(adecisions=first_decisions)) + assert await call_current_router() == "first" + + monkeypatch.setattr(proxy_server, "llm_router", SimpleNamespace(adecisions=second_decisions)) + assert await call_current_router() == "second" + + monkeypatch.setattr(proxy_server, "llm_router", None) + with pytest.raises(RuntimeError, match="router is not initialized"): + await call_current_router() diff --git a/tests/unit/proxy/lens/test_sources.py b/tests/unit/proxy/lens/test_sources.py index c95948cb90c..4ee3da5f9ab 100644 --- a/tests/unit/proxy/lens/test_sources.py +++ b/tests/unit/proxy/lens/test_sources.py @@ -4,13 +4,15 @@ from typing import Final, Literal import pytest -from litellm.proxy.lens.models import Execution, ExecutionContent, MetadataFilter, Scope, TracePart +from litellm.proxy.lens.models import Evidence, Execution, ExecutionContent, MetadataFilter, Scope, TracePart from litellm.proxy.lens.sources import SourceReader, execution_id, parse_execution from litellm.rust_bridge.trace.generated.models import ( ActivityAvailability, AgentRow, + CountRow, ExecutionRow, LensContentParams, + LensEvidenceParams, PartRow, ) from tests.unit.proxy.lens.test_state import lens @@ -181,9 +183,17 @@ async def test_recorded_times_survive_source_catalog_reads_search_and_python( class ContentStorage: async def lens_content(self, parameters: LensContentParams) -> tuple[PartRow, ...]: - assert parameters.source == source and parameters.record_team == "team" + assert ( + parameters.source == source + and parameters.record_team == "team" + and parameters.start_time == run.start_time + ) return rows + async def lens_evidence(self, parameters: LensEvidenceParams) -> tuple[CountRow, ...]: + assert parameters.start_time == run.start_time + return (CountRow(count=1),) + reader: Final = SourceReader(ContentStorage()) async def read(identity: str, cursor: str, offset: int) -> ExecutionContent: @@ -206,3 +216,8 @@ async def test_recorded_times_survive_source_catalog_reads_search_and_python( loaded: Final = await read(run.id, "", 1) assert loaded.parts == expected assert min(loaded.parts, key=lambda part: part.start_time).span_id == rows[-1].span_id + assert await reader.verify_evidence( + Scope(team_id="team"), + run, + Evidence(execution_id=run.id, span_id=rows[0].span_id, quote=rows[0].content), + ) diff --git a/tests/unit/rust_bridge/trace/test_queries.py b/tests/unit/rust_bridge/trace/test_queries.py index 3c0556b1d97..9fd2964af20 100644 --- a/tests/unit/rust_bridge/trace/test_queries.py +++ b/tests/unit/rust_bridge/trace/test_queries.py @@ -18,6 +18,7 @@ def test_named_query_rejects_offsets_outside_the_native_integer_range(offset: in "source": "traces", "id": "trace", "record_team": "team", + "start_time": "", "trace_ref": "ref", "cursor": "", "offset": offset, @@ -34,6 +35,7 @@ def test_named_query_rejects_parameters_for_a_different_query() -> None: source="traces", id="trace", record_team="team", + start_time="", trace_ref="ref", cursor="", offset=0, diff --git a/ui/litellm-dashboard/src/app/globals.css b/ui/litellm-dashboard/src/app/globals.css index 3bc500c4848..36c22097cdc 100644 --- a/ui/litellm-dashboard/src/app/globals.css +++ b/ui/litellm-dashboard/src/app/globals.css @@ -211,6 +211,10 @@ --trace-row-hover: oklch(0.975 0.008 215); --trace-row-selected: oklch(0.95 0.035 200); --trace-brand: oklch(0.6 0.13 195); + --finding-affected: var(--info); + --finding-unaffected: oklch(0.551 0.027 264.364 / 0.45); + --finding-quote: color-mix(in oklab, var(--warning) 16%, transparent); + --finding-ring: 0 0 0 1px oklch(0 0 0 / 0.06), 0 1px 2px -1px oklch(0 0 0 / 0.06), 0 2px 4px 0 oklch(0 0 0 / 0.04); --trace-border: oklch(0.92 0.01 230); --trace-line: oklch(0.88 0.03 205); --trace-card-border: oklch(0.93 0.01 230); @@ -288,6 +292,10 @@ --trace-row-hover: oklch(0.23 0.018 230); --trace-row-selected: oklch(0.29 0.05 210); --trace-brand: oklch(0.78 0.13 190); + --finding-affected: var(--info); + --finding-unaffected: oklch(0.707 0.022 261.325 / 0.35); + --finding-quote: color-mix(in oklab, var(--warning) 24%, transparent); + --finding-ring: 0 0 0 1px oklch(1 0 0 / 0.08); --trace-border: oklch(0.3 0.02 235); --trace-line: oklch(0.36 0.04 210); --trace-card-border: oklch(0.27 0.02 235); @@ -326,6 +334,10 @@ --color-trace-row-hover: var(--trace-row-hover); --color-trace-row-selected: var(--trace-row-selected); --color-trace-brand: var(--trace-brand); + --color-finding-affected: var(--finding-affected); + --color-finding-unaffected: var(--finding-unaffected); + --color-finding-quote: var(--finding-quote); + --shadow-finding-ring: var(--finding-ring); --color-trace-border: var(--trace-border); --color-trace-line: var(--trace-line); --color-trace-card-border: var(--trace-card-border); diff --git a/ui/litellm-dashboard/src/components/lens/LensWorkspace.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/LensWorkspace.integration.test.tsx index 83ff4040153..84f49ea7197 100644 --- a/ui/litellm-dashboard/src/components/lens/LensWorkspace.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/LensWorkspace.integration.test.tsx @@ -141,9 +141,7 @@ describe("Lens interactive demo", () => { await user.click(await screen.findByRole("row", { name: /Repeated lookups leave customers without an answer/ })); const finding = screen.getByRole("complementary", { name: "Finding details" }); expect(within(finding).getByText(/The support agent retries/)).toBeVisible(); - const summaries = within(finding).getAllByText("support_agent", { exact: true }); - await user.click(summaries[0]); - await user.click(within(finding).getAllByRole("button", { name: /Open original step/ })[0]); + await user.click(within(finding).getAllByRole("button", { name: "View span" })[0]); expect(await screen.findByRole("complementary", { name: "Span details" })).toHaveTextContent( "I will check that for you.", ); @@ -157,7 +155,8 @@ describe("Lens interactive demo", () => { await user.click(within(finding).getByRole("button", { name: "Back to finding" })); expect(within(finding).getByText(/The support agent retries/)).toBeVisible(); await user.click(within(finding).getByRole("button", { name: "Close finding (Esc)" })); - expect(await screen.findByRole("table", { name: "Findings" })).toBeVisible(); + expect(await screen.findByRole("grid", { name: "Findings" })).toBeVisible(); + expect(screen.queryByRole("complementary", { name: "Finding details" })).not.toBeInTheDocument(); expect(network).not.toHaveBeenCalled(); await expectUrl(onUrlUpdate, (url) => expect(url.get("demo")).toBe("true")); await expectUrl(onUrlUpdate, (url) => expect(url.has("span")).toBe(false)); @@ -292,7 +291,7 @@ describe("Lens interactive demo", () => { await expectUrl(onUrlUpdate, (url) => expect(url.get("tab")).toBe("settings")); expect(screen.queryByRole("dialog")).not.toBeInTheDocument(); const panel = within(screen.getByRole("region", { name: "Settings" })); - expect(panel.getByRole("status")).toHaveTextContent("Tracing enabled"); + expect(panel.getByText("Tracing enabled", { exact: true })).toBeVisible(); expect(panel.getByRole("heading", { name: "Analysis worker" })).toBeVisible(); expect(panel.getByRole("heading", { name: worker.name })).toBeVisible(); expect(panel.getByText("Connected")).toBeVisible(); diff --git a/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx b/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx index 1a483f55e90..6c5629d2a07 100644 --- a/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx +++ b/ui/litellm-dashboard/src/components/lens/LensWorkspace.tsx @@ -196,6 +196,7 @@ function LensContent({ userRole, readOnly }: Omit readOnly={readOnly} canMintTracingKey={isAdmin} canViewFindings={canViewInvestigations} + onSetUpSignals={canConfigure ? showSettings : undefined} /> diff --git a/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts b/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts index 422601c44c2..ddc7019eadb 100644 --- a/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts +++ b/ui/litellm-dashboard/src/components/lens/data/demo/createLensDemo.ts @@ -37,6 +37,8 @@ function demoLensApi(data: LensDemoData): LensApi { saveLens: readOnly, startRun: readOnly, watchAll: async () => ({ watching: [], skipped: [] }), + signalConfig: async () => ({ model: "", threshold: 0.5, signals: [] }), + saveSignalConfig: readOnly, cancelRun: readOnly, reviewFinding: readOnly, registerWorker: readOnly, @@ -90,6 +92,8 @@ function demoTracesApi(data: LensDemoData): TracesApi { ); return { ...trace, finding_count: assessed.length ? findings.size : null }; }), + signals: async (traces) => + traces.map((trace) => ({ ...trace, status: "unclassified" as const, flags: [], model: "", classified_at: null })), anyRecorded: async () => data.runs.length > 0, trace: (traceId) => found(run(traceId)?.trace), span: (traceId, spanId) => found(run(traceId)?.details.find((span) => span.span_id === spanId)), diff --git a/ui/litellm-dashboard/src/components/lens/data/queries.ts b/ui/litellm-dashboard/src/components/lens/data/queries.ts index a918d33107d..63c311ecd1c 100644 --- a/ui/litellm-dashboard/src/components/lens/data/queries.ts +++ b/ui/litellm-dashboard/src/components/lens/data/queries.ts @@ -25,6 +25,7 @@ export const lensKeys = { models: (scope: string) => [...lensKeys.all, "models", { scope }] as const, modelDetails: (scope: string) => [...lensKeys.all, "model-details", { scope }] as const, activity: (scope: string) => [...lensKeys.all, "activity-available", { scope }] as const, + signalConfig: (scope: string) => [...lensKeys.all, "signal-config", { scope }] as const, discoveries: () => [...lensKeys.all, "discovery"] as const, discovery: (scope: string, source: Settings["source"], hours: number | undefined) => [...lensKeys.discoveries(), { scope, source, hours }] as const, @@ -46,6 +47,13 @@ export const lensQueries = { modelDetails(api: LensApi) { return queryOptions({ queryKey: lensKeys.modelDetails(api.scope), queryFn: () => api.modelDetails() }); }, + signalConfig(api: LensApi) { + return queryOptions({ + queryKey: lensKeys.signalConfig(api.scope), + queryFn: () => api.signalConfig(), + staleTime: 5000, + }); + }, activity(api: LensApi, loaded: boolean) { const options = { queryKey: lensKeys.activity(api.scope), diff --git a/ui/litellm-dashboard/src/components/lens/data/service.ts b/ui/litellm-dashboard/src/components/lens/data/service.ts index 9de75206902..dd210a2bb31 100644 --- a/ui/litellm-dashboard/src/components/lens/data/service.ts +++ b/ui/litellm-dashboard/src/components/lens/data/service.ts @@ -13,6 +13,7 @@ import type { RunWindow, Sample, Settings, + SignalConfig, WorkerCreated, } from "../model/types"; @@ -62,6 +63,8 @@ export interface LensApi { saveLens(id: string | undefined, settings: Settings): Promise; startRun(lensId: string, request?: RunWindow): Promise; watchAll(): Promise; + signalConfig(): Promise; + saveSignalConfig(config: SignalConfig): Promise; cancelRun(lensId: string): Promise; reviewFinding(lensId: string, findingId: string, status: FindingStatus, reason: string): Promise; registerWorker(analysisKeyId: string): Promise; @@ -168,6 +171,8 @@ export function liveLensApi(client: LensClient, apiClient: ApiClient, accessToke ), startRun: (lensId, request = {}) => sent(client.POST("/lens/{lens_id}/runs", { ...lens(lensId), body: request })), watchAll: () => required(client.POST("/lens/watch-all", { headers })), + signalConfig: () => required(client.GET("/lens/signals", { headers })), + saveSignalConfig: (config) => required(client.PUT("/lens/signals", { headers, body: config })), cancelRun: (lensId) => sent(client.POST("/lens/{lens_id}/cancel", lens(lensId))), reviewFinding: (lensId, findingId, status, reason) => sent( diff --git a/ui/litellm-dashboard/src/components/lens/investigations/FindingDetails.integration.test.tsx b/ui/litellm-dashboard/src/components/lens/investigations/FindingDetails.integration.test.tsx index e705ce17e81..a6c1e60cb3b 100644 --- a/ui/litellm-dashboard/src/components/lens/investigations/FindingDetails.integration.test.tsx +++ b/ui/litellm-dashboard/src/components/lens/investigations/FindingDetails.integration.test.tsx @@ -116,8 +116,7 @@ it("stacks a quote's original step over the finding and keeps the feedback draft const panel = screen.getByRole("complementary", { name: "Finding details" }); const reason = () => within(panel).getByRole("textbox", { name: "What should Lens remember?", hidden: true }); fireEvent.change(reason(), { target: { value: "Draft feedback" } }); - for (const summary of within(panel).getAllByText(/quote$/)) await user.click(summary); - await user.click(within(panel).getAllByRole("button", { name: "Open original step" })[0]); + await user.click(within(panel).getAllByRole("button", { name: "View span" })[0]); expect(await within(panel).findByTestId("run-view")).toHaveTextContent("trace-1 at step-a"); expect(reason()).not.toBeVisible(); expect(screen.queryByRole("dialog")).not.toBeInTheDocument(); @@ -127,13 +126,51 @@ it("stacks a quote's original step over the finding and keeps the feedback draft expect(reason()).toBeVisible(); expect(reason()).toHaveValue("Draft feedback"); - await user.click(within(panel).getAllByRole("button", { name: "Open original step" })[1]); + await user.click(within(panel).getAllByRole("button", { name: "View span" })[1]); expect(await within(panel).findByTestId("run-view")).toHaveTextContent("trace-2 at step-b"); const url = new URLSearchParams(String(onUrlUpdate.mock.lastCall?.[0].queryString ?? "")); expect(url.get("evidence")).toBe(traceOf("trace-2")); expect(url.get("evidence_span")).toBe("step-b"); }); +it("reports how many sampled traces the finding affected and highlights each quoted line", () => { + const traceOf = (id: string) => btoa(JSON.stringify(["traces", "", id])); + const sampled = ["a", "b", "c", "d"].map((id) => ({ + id: traceOf(id), + name: `run ${id}`, + start_time: "2026-10-01T10:00:00Z", + metadata: [], + root_seen: true, + service: "support_agent", + source: "traces" as const, + span_count: 1, + team_id: "", + trace_id: id, + trace_ref: "", + })); + const current: Finding = { + ...finding, + occurrences: [traceOf("a")], + evidence: [{ execution_id: traceOf("a"), span_id: "s", quote: "files:read is missing", role: "support" }], + }; + renderWithLens( + + + , + ); + expect(screen.getByRole("region", { name: "Frequency" })).toHaveTextContent(/25%\s*1 of 4 traces affected/); + const example = screen.getByRole("article", { name: "run a" }); + expect(within(example).getByText("files:read is missing").tagName).toBe("MARK"); + expect(screen.queryByRole("article", { name: "run b" })).not.toBeInTheDocument(); +}); + it("shows contributing investigation runs and every affected trace, including older traces without retained quotes", async () => { const traceId = btoa(JSON.stringify(["traces", "", "older-trace", ""])); const current: Finding = { @@ -142,8 +179,33 @@ it("shows contributing investigation runs and every affected trace, including ol investigation_runs: ["first-investigation-run", "second-investigation-run"], }; renderWithLens(); - expect(screen.getByText("Found across 2 investigation runs")).toBeInTheDocument(); - expect(screen.getByText(/1 affected trace/)).toBeInTheDocument(); - fireEvent.click(screen.getByText("older-trace")); - expect(screen.getByRole("button", { name: "Open original trace" })).toBeInTheDocument(); + expect(screen.getByText("1 affected trace")).toBeVisible(); + expect(screen.getByText("Found across 2 investigation runs")).toBeVisible(); + const example = screen.getByRole("article", { name: "Trace older-tr" }); + expect(within(example).getByText("No quote was retained for this trace.")).toBeVisible(); + expect(within(example).getByRole("button", { name: "View trace" })).toBeVisible(); +}); + +it("shows the finding's priority and keeps the first three examples, revealing the rest on request", async () => { + const user = userEvent.setup(); + const traceOf = (id: string) => btoa(JSON.stringify(["traces", "", id])); + const ids = ["t1", "t2", "t3", "t4", "t5"]; + const current: Finding = { + ...finding, + occurrences: ids.map(traceOf), + evidence: ids.map((id) => ({ + execution_id: traceOf(id), + span_id: id, + quote: `Input: ${id}\nOutput: done`, + role: "support" as const, + })), + }; + renderWithLens(); + const panel = screen.getByRole("complementary", { name: "Finding details" }); + expect(within(panel).getByText("High priority")).toBeVisible(); + expect(within(panel).getAllByRole("article")).toHaveLength(3); + expect(within(panel).getAllByText("Call and result")).toHaveLength(3); + await user.click(within(panel).getByRole("button", { name: "Show 2 more examples" })); + expect(within(panel).getAllByRole("article")).toHaveLength(5); + expect(within(panel).queryByRole("button", { name: /Show \d+ more/ })).not.toBeInTheDocument(); }); diff --git a/ui/litellm-dashboard/src/components/lens/investigations/FindingDetails.tsx b/ui/litellm-dashboard/src/components/lens/investigations/FindingDetails.tsx index 95157a878d1..59ad282bbe0 100644 --- a/ui/litellm-dashboard/src/components/lens/investigations/FindingDetails.tsx +++ b/ui/litellm-dashboard/src/components/lens/investigations/FindingDetails.tsx @@ -1,23 +1,32 @@ "use client"; import { useState } from "react"; -import { ArrowUpRight } from "lucide-react"; +import { ChevronRight, ClipboardCopy, X } from "lucide-react"; import { Inspector } from "@/components/shared/Inspector"; import { Button } from "@/components/ui/button"; import { Textarea } from "@/components/ui/textarea"; +import { useNow } from "@/hooks/useNow"; +import { copyToClipboard } from "@/utils/dataUtils"; import { AddToDatasetButton } from "../datasets/AddToDatasetDialog"; -import { evidenceTarget } from "../model/findings"; -import { runTime } from "../model/format"; +import { evidenceTarget, findingMarkdown } from "../model/findings"; +import { findingFrequency } from "../model/frequency"; +import { agoLabel, runTime } from "../model/format"; import { findingAgents, findingKey, type OwnedFinding, sampledExecutions } from "../model/inbox"; import type { Finding, Sample } from "../model/types"; import { EvidenceView } from "./Evidence"; +import { FrequencyCard } from "./FrequencyCard"; import { IssueBrief } from "./IssueBrief"; +import { PriorityPill } from "./PriorityMark"; import { type EvidenceRef, useEvidenceRoute } from "../route"; export const ownedFindingKey = (owned: OwnedFinding): string => findingKey(owned.lens, owned.finding); +type Quote = Finding["evidence"][number]; + +const SECTION_LABEL = "text-xs font-medium text-muted-foreground"; + export interface FindingDetailsProps { readonly finding: Finding; readonly lensId?: string; @@ -27,6 +36,240 @@ export interface FindingDetailsProps { readonly busy: boolean; readonly onOpenEvidence: (evidence: EvidenceRef) => void; readonly onReview: (status: Finding["status"], reason: string) => void; + readonly onClose?: () => void; +} + +function TopBar({ finding, onClose }: Pick) { + const now = useNow(30000); + return ( +
+

+ + {finding.id.slice(0, 8)} + + + + {agoLabel(Date.parse(finding.last_seen), now)} + +

+
+ + {onClose && ( + + )} +
+
+ ); +} + +function Disclosure({ title, children }: { title: string; children: React.ReactNode }) { + return ( +
+ + +
{children}
+
+ ); +} + +function ProseSection({ title, children }: { title: string; children: string }) { + return ( +
+

{title}

+

+ {children} +

+
+ ); +} + +const FIELD = /^(Input|Output|Status|Error)\s*:/gm; +const FIELD_LABEL: Readonly> = { + "Input,Output": "Call and result", + Input: "Call input", + Output: "Returned output", + Status: "Span status", + Error: "Error", +}; + +function quoteLabel(quote: Quote, isTrace: boolean): string { + if (quote.role === "counterexample") return "Counterexample"; + const fields = [...new Set(Array.from(quote.quote.matchAll(FIELD), (m) => m[1]))].join(","); + return FIELD_LABEL[fields] ?? (isTrace ? "Trace step" : "Logged request"); +} + +const MARK = { + support: "rounded-sm bg-finding-quote px-0.5 text-inherit", + counterexample: "rounded-sm bg-success/20 px-0.5 text-inherit", +} as const; + +function QuoteCard({ quote, onOpen }: { quote: Quote; onOpen: () => void }) { + const isTrace = evidenceTarget(quote.execution_id)?.source === "traces"; + return ( +
+
+ {quoteLabel(quote, isTrace)} + +
+ +
+ ); +} + +function EvidenceRail({ children }: { children: React.ReactNode }) { + return ( + <> +
+
+
+
+
+ +
Evidence
+
+
+ + + ); +} + +interface ExampleGroup { + readonly id: string; + readonly run: Sample["executions"][number] | undefined; + readonly quotes: readonly Quote[]; +} + +function Example({ group, onOpenEvidence }: { group: ExampleGroup; onOpenEvidence: (e: EvidenceRef) => void }) { + const traceId = evidenceTarget(group.id)?.id; + const name = group.run?.name ?? (traceId ? `Trace ${traceId.slice(0, 8)}` : "Recorded run"); + return ( +
+
+

+ {name} +

+ + {[group.run?.service, group.run && runTime(group.run.start_time)].filter(Boolean).join(" · ")} + +
+
+ {group.quotes.length === 0 ? ( +
+

No quote was retained for this trace.

+ +
+ ) : ( + + {group.quotes.map((quote, i) => ( + onOpenEvidence({ id: quote.execution_id, span: quote.span_id })} + /> + ))} + + )} +
+
+ ); +} + +const VISIBLE_EXAMPLES = 3; + +function Examples({ + groups, + onOpenEvidence, +}: { + groups: readonly ExampleGroup[]; + onOpenEvidence: (e: EvidenceRef) => void; +}) { + const [expanded, setExpanded] = useState(false); + if (groups.length === 0) return

No examples were recorded.

; + const shown = expanded ? groups : groups.slice(0, VISIBLE_EXAMPLES); + const hidden = groups.length - shown.length; + return ( + <> + {shown.map((group) => ( + + ))} + {hidden > 0 && ( + + )} + + ); +} + +function ReviewForm({ finding, busy, onReview }: Pick) { + const [reason, setReason] = useState(finding.reason ?? ""); + return ( +
+