diff --git a/litellm-rust/crates/traces-clickhouse/query/list_traces.sql b/litellm-rust/crates/traces-clickhouse/query/list_traces.sql index 1598f04aba2..c52adf7ef49 100644 --- a/litellm-rust/crates/traces-clickhouse/query/list_traces.sql +++ b/litellm-rust/crates/traces-clickhouse/query/list_traces.sql @@ -18,8 +18,7 @@ SELECT TraceId AS trace_id, FROM agent_traces_by_key WHERE ({all_teams:UInt8} = 1 OR ({user_id:String} != '' AND UserIds = [{user_id:String}]) - OR has({team_ids:Array(String)}, TeamId) - OR ({api_key_hash:String} != '' AND ApiKeyHash = {api_key_hash:String})) + OR has({team_ids:Array(String)}, TeamId)) GROUP BY TeamId, ApiKeyHash, TraceId HAVING min(StartTs) >= fromUnixTimestamp64Milli({start_ms:Int64}) AND min(StartTs) < fromUnixTimestamp64Milli({end_ms:Int64}) diff --git a/litellm-rust/crates/traces-clickhouse/query/span_detail.sql b/litellm-rust/crates/traces-clickhouse/query/span_detail.sql index 4da48e00d4b..7db742ea3ee 100644 --- a/litellm-rust/crates/traces-clickhouse/query/span_detail.sql +++ b/litellm-rust/crates/traces-clickhouse/query/span_detail.sql @@ -9,8 +9,7 @@ LEFT JOIN ( AND ObservationType = 'llm' AND Output != '' AND ({all_teams:UInt8} = 1 OR ({user_id:String} != '' AND UserId = {user_id:String}) - OR has({team_ids:Array(String)}, TeamId) - OR ({api_key_hash:String} != '' AND ApiKeyHash = {api_key_hash:String})) + OR has({team_ids:Array(String)}, TeamId)) AND ({trace_ref:String} = '' OR hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId))) = {trace_ref:String}) GROUP BY TeamId, ApiKeyHash, ParentSpanId @@ -19,8 +18,7 @@ LEFT JOIN ( WHERE o.TraceId = {trace_id:String} AND o.SpanId = {span_id:String} AND ({all_teams:UInt8} = 1 OR ({user_id:String} != '' AND o.UserId = {user_id:String}) - OR has({team_ids:Array(String)}, o.TeamId) - OR ({api_key_hash:String} != '' AND o.ApiKeyHash = {api_key_hash:String})) + OR has({team_ids:Array(String)}, o.TeamId)) AND ({trace_ref:String} = '' OR hex(SHA256(concat(o.TeamId, char(0), o.ApiKeyHash, char(0), o.TraceId))) = {trace_ref:String}) LIMIT 1 diff --git a/litellm-rust/crates/traces-clickhouse/query/span_error.sql b/litellm-rust/crates/traces-clickhouse/query/span_error.sql index 087962227a8..e1226c4d23c 100644 --- a/litellm-rust/crates/traces-clickhouse/query/span_error.sql +++ b/litellm-rust/crates/traces-clickhouse/query/span_error.sql @@ -6,8 +6,7 @@ FROM otel_traces WHERE TraceId = {trace_id:String} AND SpanId = {span_id:String} AND ({all_teams:UInt8} = 1 OR ({user_id:String} != '' AND UserId = {user_id:String}) - OR has({team_ids:Array(String)}, TeamId) - OR ({api_key_hash:String} != '' AND ApiKeyHash = {api_key_hash:String})) + OR has({team_ids:Array(String)}, TeamId)) AND ({trace_ref:String} = '' OR hex(SHA256(concat(TeamId, char(0), ApiKeyHash, char(0), TraceId))) = {trace_ref:String}) AND ({error_version:String} = '' OR hex(SHA256(StatusMessage)) = {error_version:String}) diff --git a/litellm-rust/crates/traces-clickhouse/query/spend_by_response_ids.sql b/litellm-rust/crates/traces-clickhouse/query/spend_by_response_ids.sql index eb132099ac5..963c832232b 100644 --- a/litellm-rust/crates/traces-clickhouse/query/spend_by_response_ids.sql +++ b/litellm-rust/crates/traces-clickhouse/query/spend_by_response_ids.sql @@ -6,6 +6,5 @@ WHERE response_id IN {response_ids:Array(String)} AND start_time < fromUnixTimestamp64Milli({end_ms:Int64}) AND ({all_teams:UInt8} = 1 OR ({user_id:String} != '' AND user = {user_id:String}) - OR has({team_ids:Array(String)}, team_id) - OR ({api_key_hash:String} != '' AND api_key = {api_key_hash:String})) + OR has({team_ids:Array(String)}, team_id)) ORDER BY start_time DESC diff --git a/litellm-rust/crates/traces-clickhouse/query/trace_identity.sql b/litellm-rust/crates/traces-clickhouse/query/trace_identity.sql index 6b860065468..e3881b150b7 100644 --- a/litellm-rust/crates/traces-clickhouse/query/trace_identity.sql +++ b/litellm-rust/crates/traces-clickhouse/query/trace_identity.sql @@ -3,7 +3,6 @@ FROM otel_traces WHERE TraceId = {trace_id:String} AND ({all_teams:UInt8} = 1 OR ({user_id:String} != '' AND UserId = {user_id:String}) - OR has({team_ids:Array(String)}, TeamId) - OR ({api_key_hash:String} != '' AND ApiKeyHash = {api_key_hash:String})) + OR has({team_ids:Array(String)}, TeamId)) GROUP BY TeamId, ApiKeyHash, TraceId LIMIT 2 diff --git a/litellm-rust/crates/traces-clickhouse/query/trace_spans.sql b/litellm-rust/crates/traces-clickhouse/query/trace_spans.sql index ade1fdf0d86..f00329d8136 100644 --- a/litellm-rust/crates/traces-clickhouse/query/trace_spans.sql +++ b/litellm-rust/crates/traces-clickhouse/query/trace_spans.sql @@ -12,8 +12,7 @@ FROM otel_traces AS o WHERE o.TraceId = {trace_id:String} AND ({all_teams:UInt8} = 1 OR ({user_id:String} != '' AND o.UserId = {user_id:String}) - OR has({team_ids:Array(String)}, o.TeamId) - OR ({api_key_hash:String} != '' AND o.ApiKeyHash = {api_key_hash:String})) + OR has({team_ids:Array(String)}, o.TeamId)) AND ({trace_ref:String} = '' OR hex(SHA256(concat(o.TeamId, char(0), o.ApiKeyHash, char(0), o.TraceId))) = {trace_ref:String}) ORDER BY o.Timestamp, o.EngineReceivedMs, o.StatusMessage diff --git a/litellm-rust/crates/traces-clickhouse/src/query.rs b/litellm-rust/crates/traces-clickhouse/src/query.rs index 459f0f95fe7..9f9c5463b7e 100644 --- a/litellm-rust/crates/traces-clickhouse/src/query.rs +++ b/litellm-rust/crates/traces-clickhouse/src/query.rs @@ -425,7 +425,7 @@ pub async fn query_help(client: &Client, connection: &Connection) -> Result( - json!({"all_teams": 0, "user_id": "user", "team_ids": ["team-a", "team-b"], "api_key_hash": "key", "start_ms": -1, "end_ms": 10, "cursor_ms": 0, "cursor_trace_id": "", "limit": u32::MAX}), + json!({"all_teams": 0, "user_id": "user", "team_ids": ["team-a", "team-b"], "start_ms": -1, "end_ms": 10, "cursor_ms": 0, "cursor_trace_id": "", "limit": u32::MAX}), quoted, ); round_trip::( - json!({"all_teams": 0, "user_id": "", "team_ids": [], "api_key_hash": "key", "trace_id": "trace", "trace_ref": "ref", "span_id": "span", "error_offset": u64::MAX, "error_version": "version"}), + json!({"all_teams": 0, "user_id": "", "team_ids": [], "trace_id": "trace", "trace_ref": "ref", "span_id": "span", "error_offset": u64::MAX, "error_version": "version"}), quoted, ); round_trip::( - json!({"all_teams": 0, "user_id": "user", "team_ids": ["team-a", "team-b"], "api_key_hash": "", "response_ids": ["response"], "start_ms": -1, "end_ms": 10}), + json!({"all_teams": 0, "user_id": "user", "team_ids": ["team-a", "team-b"], "response_ids": ["response"], "start_ms": -1, "end_ms": 10}), quoted, ); } diff --git a/litellm-rust/crates/traces-clickhouse/src/query_access.rs b/litellm-rust/crates/traces-clickhouse/src/query_access.rs index 5d82a5e2e93..e6ca322098e 100644 --- a/litellm-rust/crates/traces-clickhouse/src/query_access.rs +++ b/litellm-rust/crates/traces-clickhouse/src/query_access.rs @@ -176,17 +176,13 @@ impl QueryReaders { } fn predicate(scope: &QueryScope, table: TraceTable) -> String { - let (team, key) = match table { - TraceTable::OtelTraces | TraceTable::AgentTracesByKey => ("TeamId", "ApiKeyHash"), - TraceTable::SpendLogs => ("team_id", "api_key"), + let team = match table { + TraceTable::OtelTraces | TraceTable::AgentTracesByKey => "TeamId", + TraceTable::SpendLogs => "team_id", }; match scope { - QueryScope::Admin => "1".to_owned(), - QueryScope::Logs { - user_id, - team_ids, - api_key_hash, - } => { + QueryScope::All => "1".to_owned(), + QueryScope::Owned { user_id, team_ids } => { let owner = literal(user_id); let user_clause = match table { TraceTable::OtelTraces => format!("UserId = {owner}"), @@ -203,21 +199,8 @@ fn predicate(scope: &QueryScope, table: TraceTable) -> String { } else { format!("{team} IN ({teams})") }; - format!( - "({owner} != '' AND {user_clause}) OR ({team_clause}) OR ({hash} != '' AND {key} = {hash})", - owner = owner, - hash = literal(api_key_hash), - ) + format!("({owner} != '' AND {user_clause}) OR ({team_clause})") } - QueryScope::Team { team_id } => format!("{team} = {}", literal(team_id)), - QueryScope::Key { - team_id, - api_key_hash, - } => format!( - "{team} = {} AND {key} = {}", - literal(team_id), - literal(api_key_hash) - ), } } @@ -239,33 +222,24 @@ mod tests { use rstest::rstest; #[rstest] - #[case::otel(TraceTable::OtelTraces, "TeamId", "ApiKeyHash")] - #[case::agent(TraceTable::AgentTracesByKey, "TeamId", "ApiKeyHash")] - #[case::spend(TraceTable::SpendLogs, "team_id", "api_key")] + #[case::otel(TraceTable::OtelTraces, "TeamId", "UserId = ''")] + #[case::agent(TraceTable::AgentTracesByKey, "TeamId", "UserIds = ['']")] + #[case::spend(TraceTable::SpendLogs, "team_id", "user = ''")] fn predicates_preserve_scope_and_escape_values( #[case] table: TraceTable, #[case] team: &str, - #[case] key: &str, + #[case] user: &str, ) { - assert_eq!(predicate(&QueryScope::Admin, table), "1"); + assert_eq!(predicate(&QueryScope::All, table), "1"); assert_eq!( predicate( - &QueryScope::Team { - team_id: "team'\\".into() + &QueryScope::Owned { + user_id: String::new(), + team_ids: vec!["team'\\".into()] }, table ), - format!("{team} = 'team\\'\\\\'") - ); - assert_eq!( - predicate( - &QueryScope::Key { - team_id: "".into(), - api_key_hash: "key'\\".into() - }, - table - ), - format!("{team} = '' AND {key} = 'key\\'\\\\'") + format!("('' != '' AND {user}) OR ({team} IN ('team\\'\\\\'))") ); } } diff --git a/litellm-rust/crates/traces-clickhouse/src/sql.rs b/litellm-rust/crates/traces-clickhouse/src/sql.rs index 0f739314418..16a610ca13a 100644 --- a/litellm-rust/crates/traces-clickhouse/src/sql.rs +++ b/litellm-rust/crates/traces-clickhouse/src/sql.rs @@ -64,7 +64,7 @@ mod tests { #[case] specific: serde_json::Value, ) { let common = serde_json::json!({ - "all_teams": 1, "user_id": "", "team_ids": [], "api_key_hash": "", "trace_id": "trace", "trace_ref": "" + "all_teams": 1, "user_id": "", "team_ids": [], "trace_id": "trace", "trace_ref": "" }); let parameters: BTreeMap = common .as_object() diff --git a/litellm-rust/crates/traces-clickhouse/templates/query_help.jinja b/litellm-rust/crates/traces-clickhouse/templates/query_help.jinja index 6e746221022..3bea0ff4efe 100644 --- a/litellm-rust/crates/traces-clickhouse/templates/query_help.jinja +++ b/litellm-rust/crates/traces-clickhouse/templates/query_help.jinja @@ -46,7 +46,7 @@ Gotchas {% block reader_limits %}The reader enforces {{ limits.result_rows }} result rows, {{ limits.result_mib() }} MiB response bytes, {{ limits.memory_mib() }} MiB memory and a {{ limits.execution_seconds }} second query limit; exceeding limits fails instead of returning partial results{% endblock %} -{% block reader_profile %}LiteLLM provisions SELECT-only readers from the configured ClickHouse connection and enforces request-log visibility through row policies. Callers see their own user rows and permitted teams, or their own key rows when no user identity is available. Provisioning requires CREATE USER, ALTER USER, CREATE ROW POLICY, and GRANT SELECT permissions{% endblock %} +{% block reader_profile %}LiteLLM provisions SELECT-only readers from the configured ClickHouse connection and enforces request-log visibility through row policies. Callers see their own user rows and permitted teams. Provisioning requires CREATE USER, ALTER USER, CREATE ROW POLICY, and GRANT SELECT permissions{% endblock %} {% block output_format %}Do not add FORMAT clauses; the endpoint requires ClickHouse JSON output{% endblock %} diff --git a/litellm-rust/crates/traces-clickhouse/tests/admin_sql.rs b/litellm-rust/crates/traces-clickhouse/tests/admin_sql.rs index 614dad7a35a..49499ebfa39 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/admin_sql.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/admin_sql.rs @@ -41,7 +41,7 @@ async fn database() -> Result> { } let readers = QueryReaders::new(Connection::writer(&admin_url)?, "litellm".into()); let connection = readers - .connection(&client, &QueryScope::Admin, "test-secret") + .connection(&client, &QueryScope::All, "test-secret") .await?; let url = connection.url().to_string(); Ok(Database { diff --git a/litellm-rust/crates/traces-clickhouse/tests/migrations.rs b/litellm-rust/crates/traces-clickhouse/tests/migrations.rs index 9fdeb450c70..2f9e88af3e7 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/migrations.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/migrations.rs @@ -104,7 +104,6 @@ async fn schema_supports_span_rollups_and_spend_joins( all_teams: 0, user_id: String::new(), team_ids: vec!["team-1".into()], - api_key_hash: String::new(), }, trace_id: "trace-1".into(), trace_ref: String::new(), @@ -119,7 +118,6 @@ async fn schema_supports_span_rollups_and_spend_joins( ("all_teams".into(), Parameter::Integer(0)), ("user_id".into(), Parameter::Text(String::new())), ("team_ids".into(), Parameter::Strings(vec!["team-1".into()])), - ("api_key_hash".into(), Parameter::Text(String::new())), ( "start_ms".into(), Parameter::Integer(timestamp / 1_000_000 - 1000), @@ -153,7 +151,6 @@ async fn schema_supports_span_rollups_and_spend_joins( ("all_teams".into(), Parameter::Integer(0)), ("user_id".into(), Parameter::Text(String::new())), ("team_ids".into(), Parameter::Strings(vec!["team-1".into()])), - ("api_key_hash".into(), Parameter::Text(String::new())), ( "start_ms".into(), Parameter::Integer(timestamp / 1_000_000 - 1000), @@ -421,6 +418,7 @@ async fn listed_agent_names_preserve_scope_and_cursor( vec![serde_json::from_value(serde_json::json!({ "Timestamp": timestamp, "TraceId": trace, "SpanId": span, "ParentSpanId": parent, "ServiceName": "shared-app", "SpanName": span, "AgentName": agent, + "UserId": if key == "one" { "owner" } else { "other" }, "Framework": framework, "ObservationType": "agent", "ResourceAttributes": {"litellm.team_id": team, "litellm.api_key_hash": key} }))?], @@ -442,9 +440,8 @@ async fn listed_agent_names_preserve_scope_and_cursor( let connection = Connection::configured(&database.url, "trace_test", "default", "")?; let parameters = BTreeMap::from([ ("all_teams".into(), Parameter::Integer(0)), - ("user_id".into(), Parameter::Text(String::new())), + ("user_id".into(), Parameter::Text("owner".into())), ("team_ids".into(), Parameter::Strings(vec![])), - ("api_key_hash".into(), Parameter::Text("one".into())), ( "start_ms".into(), Parameter::Integer(timestamp / 1_000_000 - 1000), @@ -596,7 +593,6 @@ async fn rollup_merges_spans_across_days_without_losing_root_fields( ("all_teams".into(), Parameter::Integer(0)), ("user_id".into(), Parameter::Text(String::new())), ("team_ids".into(), Parameter::Strings(vec!["team-1".into()])), - ("api_key_hash".into(), Parameter::Text(String::new())), ( "start_ms".into(), Parameter::Integer(day_start / 1_000_000 - 2000), @@ -795,7 +791,7 @@ async fn lens_filters_reads_and_evidence_keep_reused_trace_ids_separate( for (key, text) in [("one", "timeout"), ("two", "success")] { insert_rows(&database, "otel_traces", vec![serde_json::from_value(serde_json::json!({ "Timestamp": timestamp, "TraceId": "shared", "SpanId": "root", "ParentSpanId": "", - "ServiceName": "review", "SpanName": "release", "Input": text, + "ServiceName": "review", "SpanName": "release", "Input": text, "UserId": key, "ResourceAttributes": {"litellm.team_id": "team", "litellm.api_key_hash": key, "swarm": "release"} }))?]).await?; } @@ -849,7 +845,6 @@ async fn lens_filters_reads_and_evidence_keep_reused_trace_ids_separate( ("all_teams".into(), Parameter::Integer(0)), ("user_id".into(), Parameter::Text(String::new())), ("team_ids".into(), Parameter::Strings(vec!["team".into()])), - ("api_key_hash".into(), Parameter::Text(String::new())), ]); let identities: serde_json::Value = serde_json::from_str( &execute_named_read( @@ -861,11 +856,11 @@ async fn lens_filters_reads_and_evidence_keep_reused_trace_ids_separate( .await?, )?; assert_eq!(identities["data"].as_array().map(Vec::len), Some(2)); - let key_params = identity_params + let user_params = identity_params .into_iter() .chain([ ("team_ids".into(), Parameter::Strings(vec![])), - ("api_key_hash".into(), Parameter::Text("one".into())), + ("user_id".into(), Parameter::Text("one".into())), ]) .collect(); let identity: serde_json::Value = serde_json::from_str( @@ -873,7 +868,7 @@ async fn lens_filters_reads_and_evidence_keep_reused_trace_ids_separate( &database.client, &connection, ReadQuery::TraceIdentity, - &key_params, + &user_params, ) .await?, )?; @@ -1170,7 +1165,6 @@ async fn trace_error_previews_preserve_paginated_diagnostics( ("all_teams".into(), Parameter::Integer(1)), ("user_id".into(), Parameter::Text(String::new())), ("team_ids".into(), Parameter::Strings(vec![])), - ("api_key_hash".into(), Parameter::Text(String::new())), ("trace_ref".into(), Parameter::Text(String::new())), ]); let body = execute_named_read( @@ -1217,10 +1211,7 @@ async fn trace_error_previews_preserve_paginated_diagnostics( } assert_eq!(recovered, message); parameters.insert("all_teams".into(), Parameter::Integer(0)); - parameters.insert( - "api_key_hash".into(), - Parameter::Text("unrelated-key".into()), - ); + parameters.insert("user_id".into(), Parameter::Text("unrelated-user".into())); let denied = execute_named_read(&database.client, &reader, ReadQuery::SpanError, ¶meters).await?; assert_eq!( @@ -1265,7 +1256,6 @@ async fn duplicate_span_preview_matches_diagnostic( ("all_teams".into(), Parameter::Integer(1)), ("user_id".into(), Parameter::Text(String::new())), ("team_ids".into(), Parameter::Strings(vec![])), - ("api_key_hash".into(), Parameter::Text(String::new())), ("trace_ref".into(), Parameter::Text(String::new())), ("error_version".into(), Parameter::Text(String::new())), ("error_offset".into(), Parameter::Integer(0)), @@ -1716,16 +1706,16 @@ fn field_definitions_match_serialized_normalized_span() { } #[rstest] -#[case::own_user("owner", vec![], "", vec!["own"])] -#[case::own_user_and_permitted_team("owner", vec!["permitted"], "", vec!["own", "team"])] -#[case::key_only("", vec![], "request-key", vec!["own"])] -#[case::no_identity("", vec![], "", vec![])] +#[case::own_user("owner", vec![], None, vec!["own"])] +#[case::own_user_and_permitted_team("owner", vec!["permitted"], None, vec!["own", "team"])] +#[case::no_identity("", vec![], None, vec![])] +#[case::legacy_key_without_identity("", vec![], Some("request-key"), vec![])] #[tokio::test] async fn named_and_sql_readers_share_request_log_visibility( #[future(awt)] database: TestResult, #[case] user: &str, #[case] teams: Vec<&str>, - #[case] key: &str, + #[case] legacy_key: Option<&str>, #[case] expected: Vec<&str>, ) -> TestResult { use litellm_traces_clickhouse::query::named::{ @@ -1748,12 +1738,10 @@ async fn named_and_sql_readers_share_request_log_visibility( let reader = Connection::reader(&database.url, "trace_test")?; let params = SpendByResponseIdsParams::from(litellm_traces::query::named::SpendByResponseIdsParams { - access: ReadAccessParams { - all_teams: 0, - user_id: user.into(), - team_ids: teams.iter().map(|team| (*team).into()).collect(), - api_key_hash: key.into(), - }, + access: serde_json::from_value::(serde_json::json!({ + "all_teams": 0, "user_id": user, "team_ids": teams, + "api_key_hash": legacy_key.unwrap_or_default(), + }))?, response_ids: vec!["shared-response".into()], start_ms: timestamp / 1_000_000 - 1, end_ms: timestamp / 1_000_000 + 1, @@ -1765,10 +1753,9 @@ async fn named_and_sql_readers_share_request_log_visibility( spend.iter().map(|row| row.0.request_id.as_str()).collect(); let expected: std::collections::BTreeSet<_> = expected.into_iter().collect(); assert_eq!(actual, expected); - let scope = QueryScope::Logs { + let scope = QueryScope::Owned { user_id: user.into(), team_ids: teams.into_iter().map(str::to_owned).collect(), - api_key_hash: key.into(), }; if user.is_empty() && scope.validate().is_err() { assert!( @@ -1829,7 +1816,6 @@ async fn rollup_cost_completeness_preserves_missing_ids_and_fails_closed_for_his all_teams: 0, user_id: "".into(), team_ids: vec!["team".into()], - api_key_hash: "".into(), }, start_ms: timestamp / 1_000_000 - 1, end_ms: timestamp / 1_000_000 + 1, @@ -1850,7 +1836,6 @@ async fn rollup_cost_completeness_preserves_missing_ids_and_fails_closed_for_his user_id: "owner".into(), team_ids: vec![], all_teams: 0, - api_key_hash: String::new(), }, ..params.0 }, @@ -1882,18 +1867,16 @@ async fn rollup_cost_completeness_preserves_missing_ids_and_fails_closed_for_his } #[rstest] -#[case::admin(1, "", vec![], "", "own answer")] -#[case::user(0, "owner", vec![], "", "own answer")] -#[case::team(0, "", vec!["alpha"], "", "own answer")] -#[case::key(0, "", vec![], "one", "own answer")] -#[case::no_identity(0, "", vec![], "", "")] +#[case::admin(1, "", vec![], "own answer")] +#[case::user(0, "owner", vec![], "own answer")] +#[case::team(0, "", vec!["alpha"], "own answer")] +#[case::no_identity(0, "", vec![], "")] #[tokio::test] async fn agent_final_answer_preserves_visibility_and_trace_ownership( #[future(awt)] database: TestResult, #[case] all_teams: u8, #[case] user: &str, #[case] teams: Vec<&str>, - #[case] key: &str, #[case] expected: &str, ) -> TestResult { use litellm_traces_clickhouse::query::named::{ReadAccessParams, SpanDetail, SpanDetailParams}; @@ -1952,7 +1935,6 @@ async fn agent_final_answer_preserves_visibility_and_trace_ownership( all_teams, user_id: user.into(), team_ids: teams.into_iter().map(str::to_owned).collect(), - api_key_hash: key.into(), }, trace_id: "shared".into(), trace_ref: String::new(), diff --git a/litellm-rust/crates/traces-clickhouse/tests/queries.rs b/litellm-rust/crates/traces-clickhouse/tests/queries.rs index 3b9f108af94..c72255d5e0d 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/queries.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/queries.rs @@ -23,23 +23,20 @@ use support::TestResult; enum ScopeCase { Admin, Team, - Key, OtherTeam, } impl ScopeCase { fn scope(self) -> QueryScope { match self { - Self::Admin => QueryScope::Admin, - Self::Team => QueryScope::Team { - team_id: "team-a".into(), + Self::Admin => QueryScope::All, + Self::Team => QueryScope::Owned { + user_id: String::new(), + team_ids: vec!["team-a".into()], }, - Self::Key => QueryScope::Key { - team_id: "team-a".into(), - api_key_hash: "key-a".into(), - }, - Self::OtherTeam => QueryScope::Team { - team_id: "team-b".into(), + Self::OtherTeam => QueryScope::Owned { + user_id: String::new(), + team_ids: vec!["team-b".into()], }, } } @@ -60,13 +57,7 @@ async fn curated_queries_return_expected_rows( #[future(awt)] seeded_database: TestResult, #[case] sql: &str, #[case] expected_json: &str, - #[values( - ScopeCase::Admin, - ScopeCase::Team, - ScopeCase::Key, - ScopeCase::OtherTeam - )] - scope: ScopeCase, + #[values(ScopeCase::Admin, ScopeCase::Team, ScopeCase::OtherTeam)] scope: ScopeCase, ) -> TestResult { let fixture = seeded_database?; let reader = fixture @@ -103,11 +94,7 @@ async fn typed_queries_read_normalized_spans_and_keep_trace_identities_separate( let fixture = seeded_database?; let reader = fixture .readers - .connection( - &fixture.database.client, - &QueryScope::Admin, - "fixture-secret", - ) + .connection(&fixture.database.client, &QueryScope::All, "fixture-secret") .await?; let params = ListTracesParams::from(contracts::ListTracesParams { access: admin_access?, @@ -176,11 +163,7 @@ async fn typed_trace_cursor_returns_the_next_fixture_trace( let fixture = seeded_database?; let reader = fixture .readers - .connection( - &fixture.database.client, - &QueryScope::Admin, - "fixture-secret", - ) + .connection(&fixture.database.client, &QueryScope::All, "fixture-secret") .await?; let params = ListTracesParams::from(contracts::ListTracesParams { access: admin_access?, @@ -218,11 +201,7 @@ async fn captured_deeplite_exports_round_trip_through_clickhouse( let decoded = insert_export(&fixture, export, "team-a", "key-a").await?; let reader = fixture .readers - .connection( - &fixture.database.client, - &QueryScope::Admin, - "fixture-secret", - ) + .connection(&fixture.database.client, &QueryScope::All, "fixture-secret") .await?; let params = TraceSpansParams { access: admin_access?, diff --git a/litellm-rust/crates/traces-clickhouse/tests/queries/read_access.json b/litellm-rust/crates/traces-clickhouse/tests/queries/read_access.json index a743b382c20..f0af446092e 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/queries/read_access.json +++ b/litellm-rust/crates/traces-clickhouse/tests/queries/read_access.json @@ -1,6 +1,8 @@ { "all_teams": 1, "user_id": "", - "team_ids": ["team-a", "team-b"], - "api_key_hash": "" + "team_ids": [ + "team-a", + "team-b" + ] } diff --git a/litellm-rust/crates/traces-clickhouse/tests/query_access.rs b/litellm-rust/crates/traces-clickhouse/tests/query_access.rs index be749351a72..ef5b76c2097 100644 --- a/litellm-rust/crates/traces-clickhouse/tests/query_access.rs +++ b/litellm-rust/crates/traces-clickhouse/tests/query_access.rs @@ -25,8 +25,8 @@ async fn database() -> Result> { let writer = Connection::parse(&url)?; ensure_schema(&client, &writer, "trace_test", 7).await?; for sql in [ - "INSERT INTO trace_test.otel_traces (TeamId, ApiKeyHash, TraceId, SpanId, Timestamp, SpanAttributes, UserId) VALUES ('team-a', 'key-a1', 'shared-trace', 'a1', now(), map('visible', 'a'), 'owner'), ('team-a', 'key-a2', 'shared-trace', 'a2', now(), map('visible', 'a'), 'other'), ('team-b', 'key-b', 'shared-trace', 'b', now(), map('secret-b', 'b'), 'owner'), ('', 'key-teamless', 'shared-trace', 'teamless', now(), map('visible', 'teamless'), ''), ('', 'key-other', 'shared-trace', 'other-teamless', now(), map('visible', 'other'), '')", - "INSERT INTO trace_test.spend_logs (team_id, api_key, request_id, start_time, end_time, metadata, user) VALUES ('team-a', 'key-a1', 'a1', now(), now(), '{\"visible\":1}', 'owner'), ('team-a', 'key-a2', 'a2', now(), now(), '{\"visible\":1}', 'other'), ('team-b', 'key-b', 'b', now(), now(), '{\"secret_b\":1}', 'owner'), ('', 'key-teamless', 'teamless', now(), now(), '{}', ''), ('', 'key-other', 'other-teamless', now(), now(), '{}', '')", + "INSERT INTO trace_test.otel_traces (TeamId, ApiKeyHash, TraceId, SpanId, Timestamp, SpanAttributes, UserId) VALUES ('team-a', 'key-a1', 'shared-trace', 'a1', now(), map('visible', 'a'), 'owner'), ('team-a', 'key-a2', 'shared-trace', 'a2', now(), map('visible', 'a'), 'other'), ('team-b', 'key-b', 'shared-trace', 'b', now(), map('secret-b', 'b'), 'owner'), ('team-c', 'key-a1', 'shared-trace', 'same-key-foreign', now(), map('visible', 'foreign'), 'other'), ('', 'key-teamless', 'shared-trace', 'teamless', now(), map('visible', 'teamless'), ''), ('', 'key-other', 'shared-trace', 'other-teamless', now(), map('visible', 'other'), '')", + "INSERT INTO trace_test.spend_logs (team_id, api_key, request_id, start_time, end_time, metadata, user) VALUES ('team-a', 'key-a1', 'a1', now(), now(), '{\"visible\":1}', 'owner'), ('team-a', 'key-a2', 'a2', now(), now(), '{\"visible\":1}', 'other'), ('team-b', 'key-b', 'b', now(), now(), '{\"secret_b\":1}', 'owner'), ('team-c', 'key-a1', 'same-key-foreign', now(), now(), '{}', 'other'), ('', 'key-teamless', 'teamless', now(), now(), '{}', ''), ('', 'key-other', 'other-teamless', now(), now(), '{}', '')", "CREATE TABLE trace_test.private_data (secret String) ENGINE = Memory", "INSERT INTO trace_test.private_data VALUES ('hidden')", ] { @@ -43,15 +43,12 @@ async fn database() -> Result> { } #[rstest] -#[case::own_user(QueryScope::Logs { user_id: "owner".into(), team_ids: vec![], api_key_hash: "".into() }, vec!["a1", "b"])] -#[case::own_user_and_permitted_team(QueryScope::Logs { user_id: "owner".into(), team_ids: vec!["team-a".into()], api_key_hash: "".into() }, vec!["a1", "a2", "b"])] -#[case::key_only_logs(QueryScope::Logs { user_id: "".into(), team_ids: vec![], api_key_hash: "key-teamless".into() }, vec!["teamless"])] -#[case::quoted_user(QueryScope::Logs { user_id: "owner' OR 1=1 --".into(), team_ids: vec![], api_key_hash: "".into() }, vec![])] -#[case::team(QueryScope::Team { team_id: "team-a".to_owned() }, vec!["a1", "a2"])] -#[case::project_key(QueryScope::Key { team_id: "team-a".to_owned(), api_key_hash: "key-a1".to_owned() }, vec!["a1"])] -#[case::teamless_key(QueryScope::Key { team_id: "".to_owned(), api_key_hash: "key-teamless".to_owned() }, vec!["teamless"])] -#[case::admin(QueryScope::Admin, vec!["a1", "a2", "b", "other-teamless", "teamless"])] -#[case::quoted_team(QueryScope::Team { team_id: "team-a' OR 1=1 --\\".to_owned() }, vec![])] +#[case::own_user(QueryScope::Owned { user_id: "owner".into(), team_ids: vec![] }, vec!["a1", "b"])] +#[case::own_user_and_permitted_team(QueryScope::Owned { user_id: "owner".into(), team_ids: vec!["team-a".into()] }, vec!["a1", "a2", "b"])] +#[case::quoted_user(QueryScope::Owned { user_id: "owner' OR 1=1 --".into(), team_ids: vec![] }, vec![])] +#[case::team(QueryScope::Owned { user_id: String::new(), team_ids: vec!["team-a".to_owned() ] }, vec!["a1", "a2"])] +#[case::admin(QueryScope::All, vec!["a1", "a2", "b", "other-teamless", "same-key-foreign", "teamless"])] +#[case::quoted_team(QueryScope::Owned { user_id: String::new(), team_ids: vec!["team-a' OR 1=1 --\\".to_owned() ] }, vec![])] #[tokio::test] async fn queries_and_help_are_scoped_by_the_database( #[future(awt)] database: Result>, @@ -111,8 +108,9 @@ async fn rotating_master_secret_revokes_previous_reader_credentials( #[future(awt)] database: Result>, ) -> Result<(), Box> { let database = database?; - let scope = QueryScope::Team { - team_id: "team-a".to_owned(), + let scope = QueryScope::Owned { + user_id: String::new(), + team_ids: vec!["team-a".to_owned()], }; let old_reader = database .readers @@ -159,8 +157,9 @@ async fn managed_reader_rejects_privilege_and_scope_bypasses( #[future(awt)] database: Result>, ) -> Result<(), Box> { let database = database?; - let scope = QueryScope::Team { - team_id: "team-a".to_owned(), + let scope = QueryScope::Owned { + user_id: String::new(), + team_ids: vec!["team-a".to_owned()], }; let reader = database .readers @@ -213,14 +212,15 @@ async fn provisioning_failure_never_returns_a_writer_connection( let database = database?; let reader = database .readers - .connection(&database.client, &QueryScope::Admin, "test-master-secret") + .connection(&database.client, &QueryScope::All, "test-master-secret") .await?; let no_provision_privileges = QueryReaders::new(reader, "trace_test".to_owned()); let result = no_provision_privileges .connection( &database.client, - &QueryScope::Team { - team_id: "team-a".to_owned(), + &QueryScope::Owned { + user_id: String::new(), + team_ids: vec!["team-a".to_owned()], }, "other-secret", ) @@ -232,7 +232,7 @@ async fn provisioning_failure_never_returns_a_writer_connection( assert!(matches!( database .readers - .connection(&database.client, &QueryScope::Admin, "") + .connection(&database.client, &QueryScope::All, "") .await, Err(Error::MissingSecret) )); @@ -241,8 +241,9 @@ async fn provisioning_failure_never_returns_a_writer_connection( .readers .connection( &database.client, - &QueryScope::Team { - team_id: String::new() + &QueryScope::Owned { + user_id: String::new(), + team_ids: vec![String::new()] }, "test-master-secret" ) @@ -263,6 +264,6 @@ async fn provisioning_failure_never_returns_a_writer_connection( ) .await?; let rows: Value = serde_json::from_str(&rows)?; - assert_eq!(rows["data"][0]["count"], 5); + assert_eq!(rows["data"][0]["count"], 6); Ok(()) } diff --git a/litellm-rust/crates/traces/src/query/named.rs b/litellm-rust/crates/traces/src/query/named.rs index b39c44d49cb..b45c076c3ba 100644 --- a/litellm-rust/crates/traces/src/query/named.rs +++ b/litellm-rust/crates/traces/src/query/named.rs @@ -6,7 +6,6 @@ pub struct ReadAccessParams { pub all_teams: u8, pub user_id: String, pub team_ids: Vec, - pub api_key_hash: String, } #[derive(Debug, Deserialize, Serialize)] diff --git a/litellm-rust/crates/traces/src/query_access.rs b/litellm-rust/crates/traces/src/query_access.rs index 51bb5c6a097..2f57bd4c0ce 100644 --- a/litellm-rust/crates/traces/src/query_access.rs +++ b/litellm-rust/crates/traces/src/query_access.rs @@ -5,33 +5,20 @@ use crate::InvalidScope; #[derive(Clone, Debug, Deserialize, Serialize)] #[serde(tag = "kind", rename_all = "snake_case", deny_unknown_fields)] pub enum QueryScope { - Admin, - Team { - team_id: String, - }, - Logs { + All, + Owned { user_id: String, team_ids: Vec, - api_key_hash: String, - }, - Key { - team_id: String, - api_key_hash: String, }, } impl QueryScope { pub fn validate(&self) -> Result<(), InvalidScope> { match self { - Self::Admin => Ok(()), - Self::Team { team_id } if !team_id.is_empty() => Ok(()), - Self::Key { api_key_hash, .. } if !api_key_hash.is_empty() => Ok(()), - Self::Logs { - user_id, - team_ids, - api_key_hash, - } if (!user_id.is_empty() || !team_ids.is_empty() || !api_key_hash.is_empty()) - && team_ids.iter().all(|team| !team.is_empty()) => + Self::All => Ok(()), + Self::Owned { user_id, team_ids } + if (!user_id.is_empty() || !team_ids.is_empty()) + && team_ids.iter().all(|team| !team.is_empty()) => { Ok(()) } diff --git a/litellm-rust/crates/traces/tests/query/named.rs b/litellm-rust/crates/traces/tests/query/named.rs index 7db77df8a29..bfe50a8684d 100644 --- a/litellm-rust/crates/traces/tests/query/named.rs +++ b/litellm-rust/crates/traces/tests/query/named.rs @@ -9,12 +9,16 @@ fn round_trip(wire: Value) { } #[rstest] -#[case::admin(vec![], "")] -#[case::multiple_teams(vec!["team-a", "team-b"], "")] -#[case::key(vec!["team-a"], "key")] -#[case::teamless_key(vec![], "key")] -fn named_requests_preserve_all_access_cases(#[case] teams: Vec<&str>, #[case] key: &str) { - let access = json!({"all_teams": u8::from(teams.is_empty() && key.is_empty()), "user_id": "", "team_ids": teams, "api_key_hash": key}); +#[case::admin(1, "", vec![])] +#[case::own_user(0, "user", vec![])] +#[case::multiple_teams(0, "user", vec!["team-a", "team-b"])] +#[case::no_identity(0, "", vec![])] +fn named_requests_preserve_all_access_cases( + #[case] all_teams: u8, + #[case] user: &str, + #[case] teams: Vec<&str>, +) { + let access = json!({"all_teams": all_teams, "user_id": user, "team_ids": teams}); round_trip::(access.clone()); let request = |specific: Value| { Value::Object( diff --git a/litellm-rust/crates/traces/tests/query_access.rs b/litellm-rust/crates/traces/tests/query_access.rs index e21abc6f427..79155065605 100644 --- a/litellm-rust/crates/traces/tests/query_access.rs +++ b/litellm-rust/crates/traces/tests/query_access.rs @@ -3,18 +3,12 @@ use rstest::rstest; use serde_json::{Value, json}; #[rstest] -#[case::admin(json!({"kind": "admin"}), true)] -#[case::team(json!({"kind": "team", "team_id": "team"}), true)] -#[case::empty_team(json!({"kind": "team", "team_id": ""}), false)] -#[case::key(json!({"kind": "key", "team_id": "team", "api_key_hash": "key"}), true)] -#[case::teamless_key(json!({"kind": "key", "team_id": "", "api_key_hash": "key"}), true)] -#[case::empty_key(json!({"kind": "key", "team_id": "team", "api_key_hash": ""}), false)] -#[case::user_logs(json!({"kind": "logs", "user_id": "user", "team_ids": [], "api_key_hash": ""}), true)] -#[case::permitted_teams(json!({"kind": "logs", "user_id": "", "team_ids": ["team"], "api_key_hash": ""}), true)] -#[case::key_logs(json!({"kind": "logs", "user_id": "", "team_ids": [], "api_key_hash": "key"}), true)] -#[case::anonymous_logs(json!({"kind": "logs", "user_id": "", "team_ids": [], "api_key_hash": ""}), false)] -#[case::empty_permitted_team(json!({"kind": "logs", "user_id": "user", "team_ids": [""], "api_key_hash": ""}), false)] -#[case::empty_teamless_key(json!({"kind": "key", "team_id": "", "api_key_hash": ""}), false)] +#[case::all(json!({"kind": "all"}), true)] +#[case::own_user(json!({"kind": "owned", "user_id": "user", "team_ids": []}), true)] +#[case::permitted_teams(json!({"kind": "owned", "user_id": "", "team_ids": ["team"]}), true)] +#[case::own_user_and_permitted_teams(json!({"kind": "owned", "user_id": "user", "team_ids": ["team"]}), true)] +#[case::no_identity(json!({"kind": "owned", "user_id": "", "team_ids": []}), false)] +#[case::empty_permitted_team(json!({"kind": "owned", "user_id": "user", "team_ids": [""]}), false)] fn scope_validation_preserves_authorization_and_wire_shape( #[case] wire: Value, #[case] valid: bool, @@ -28,22 +22,22 @@ fn scope_validation_preserves_authorization_and_wire_shape( } #[rstest] -#[case::unknown_kind(json!({"kind": "all"}))] -#[case::unknown_field(json!({"kind": "team", "team_id": "team", "extra": true}))] -#[case::missing_team(json!({"kind": "key", "api_key_hash": "key"}))] -#[case::missing_key(json!({"kind": "key", "team_id": "team"}))] +#[case::unknown_kind(json!({"kind": "unknown"}))] +#[case::unknown_field(json!({"kind": "owned", "user_id": "user", "team_ids": [], "extra": true}))] +#[case::legacy_admin(json!({"kind": "admin"}))] +#[case::legacy_logs(json!({"kind": "logs", "user_id": "user", "team_ids": []}))] +#[case::legacy_team(json!({"kind": "team", "team_id": "team"}))] +#[case::key_scope(json!({"kind": "key", "team_id": "team", "api_key_hash": "key"}))] +#[case::key_grant(json!({"kind": "owned", "user_id": "user", "team_ids": [], "api_key_hash": "key"}))] fn scope_rejects_invalid_wire_shape(#[case] wire: Value) { assert!(serde_json::from_value::(wire).is_err()); } #[rstest] -fn admin_preserves_existing_extra_field_handling() { +fn all_preserves_existing_extra_field_handling() { let scope: QueryScope = - serde_json::from_value(json!({"kind": "admin", "team_id": "ignored"})).unwrap(); - assert!(matches!(scope, QueryScope::Admin)); + serde_json::from_value(json!({"kind": "all", "team_id": "ignored"})).unwrap(); + assert!(matches!(scope, QueryScope::All)); assert!(scope.validate().is_ok()); - assert_eq!( - serde_json::to_value(scope).unwrap(), - json!({"kind": "admin"}) - ); + assert_eq!(serde_json::to_value(scope).unwrap(), json!({"kind": "all"})); } diff --git a/litellm/proxy/auth/authorization.py b/litellm/proxy/auth/authorization.py new file mode 100644 index 00000000000..91549f19d2e --- /dev/null +++ b/litellm/proxy/auth/authorization.py @@ -0,0 +1,77 @@ +from collections.abc import Awaitable, Callable, Iterable, Sequence +from dataclasses import dataclass +from typing import Final, TypeAlias + +from litellm.proxy._types import KeyManagementRoutes, LiteLLM_TeamTable, LitellmUserRoles, UserAPIKeyAuth + + +@dataclass(frozen=True, slots=True) +class AllRows: + """Unrestricted reads, granted by the consuming endpoint's role checks.""" + + +@dataclass(frozen=True, slots=True) +class OwnedRows: + """Rows owned by ``user_id`` or by any of ``team_ids``; a ``None`` user grants no own-user rows.""" + + user_id: str | None + team_ids: tuple[str, ...] = () + + +ReadScope: TypeAlias = AllRows | OwnedRows + + +async def resolve_owned_read_scope( + user_id: str | None, + permitted_team_lookup: Callable[[], Awaitable[Sequence[str]]], +) -> OwnedRows: + """Resolve own-user and permitted-team reads, falling back to own-user on lookup failure.""" + if user_id is None: + return OwnedRows(None) + try: + team_ids: Final = tuple(await permitted_team_lookup()) + except Exception: # noqa: BLE001 # preserve spend-log own-user fallback for every permission lookup failure + return OwnedRows(user_id) + return OwnedRows(user_id, team_ids) + + +def can_read_team_logs(auth: UserAPIKeyAuth, team: LiteLLM_TeamTable) -> bool: + from litellm.proxy.management.teams.access import is_team_admin + from litellm.proxy.management_endpoints.common_utils import ( + _team_member_has_permission, # pyright: ignore[reportPrivateUsage] # reuse existing team permission policy + ) + + return is_team_admin(user_api_key_dict=auth, team_obj=team) or _team_member_has_permission( + user_api_key_dict=auth, + team_obj=team, + permission=KeyManagementRoutes.SPEND_LOGS.value, + ) + + +def permitted_log_team_ids(auth: UserAPIKeyAuth, teams: Iterable[LiteLLM_TeamTable]) -> tuple[str, ...]: + return tuple(team.team_id for team in teams if can_read_team_logs(auth, team)) + + +async def can_read_log_owner( + user_id: str | None, + owner_user: str | None, + owner_team_id: str | None, + team_permission_lookup: Callable[[str], Awaitable[bool]], +) -> bool: + """Authorize stored ownership without swallowing direct team-lookup failures.""" + if owner_user is not None and owner_user == user_id: + return True + if owner_team_id: + return await team_permission_lookup(owner_team_id) + return False + + +async def resolve_trace_read_scope( + auth: UserAPIKeyAuth, + permitted_team_lookup: Callable[[], Awaitable[Sequence[str]]], +) -> ReadScope | None: + if auth.user_role in (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY): + return AllRows() + if not auth.user_id: + return None + return await resolve_owned_read_scope(auth.user_id, permitted_team_lookup) diff --git a/litellm/proxy/auth/authorization_dependencies.py b/litellm/proxy/auth/authorization_dependencies.py new file mode 100644 index 00000000000..3e7ae75dc86 --- /dev/null +++ b/litellm/proxy/auth/authorization_dependencies.py @@ -0,0 +1,57 @@ +from __future__ import annotations + +from collections.abc import Awaitable, Callable +from functools import partial +from typing import TYPE_CHECKING, Annotated, Final, TypeAlias + +from fastapi import Depends + +from litellm.proxy._types import LiteLLM_TeamTable, UserAPIKeyAuth +from litellm.proxy.auth.authorization import permitted_log_team_ids + +if TYPE_CHECKING: + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + from litellm.proxy.utils import PrismaClient, ProxyLogging + + +LogTeamLookup: TypeAlias = Callable[[UserAPIKeyAuth], Awaitable[tuple[str, ...]]] + + +async def load_permitted_log_team_ids( + auth: UserAPIKeyAuth, + *, + prisma_client: PrismaClient | None, + user_api_key_cache: UserApiKeyCache, + proxy_logging_obj: ProxyLogging, +) -> tuple[str, ...]: + from litellm.proxy.auth.auth_checks import get_user_object + from litellm.repositories.team_repository import TeamRepository + + if prisma_client is None: + return () + user_obj: Final = await get_user_object( + user_id=auth.user_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + user_id_upsert=False, + proxy_logging_obj=proxy_logging_obj, + ) + if user_obj is None or not user_obj.teams: + return () + team_rows: Final = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": user_obj.teams}}) + return permitted_log_team_ids(auth, (LiteLLM_TeamTable.model_validate(row.model_dump()) for row in team_rows)) + + +async def get_log_team_lookup() -> LogTeamLookup: + """Bind infrastructure without performing permission I/O before the handler's checks.""" + from litellm.proxy.proxy_server import prisma_client, proxy_logging_obj, user_api_key_cache + + return partial( + load_permitted_log_team_ids, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + proxy_logging_obj=proxy_logging_obj, + ) + + +LogTeamLookupDependency: TypeAlias = Annotated[LogTeamLookup, Depends(get_log_team_lookup)] diff --git a/litellm/proxy/management_endpoints/management_v1/spend_logs.py b/litellm/proxy/management_endpoints/management_v1/spend_logs.py index 1cbc454ca5e..e9e6e05ce15 100644 --- a/litellm/proxy/management_endpoints/management_v1/spend_logs.py +++ b/litellm/proxy/management_endpoints/management_v1/spend_logs.py @@ -1,12 +1,15 @@ """`/management/v1/spend_logs` facets.""" from datetime import datetime, timezone +from functools import partial from typing import Annotated, Final, Literal from fastapi import APIRouter, Depends, Query, Request from litellm._logging import verbose_proxy_logger from litellm.proxy._types import CommonProxyErrors, UserAPIKeyAuth +from litellm.proxy.auth.authorization import resolve_owned_read_scope +from litellm.proxy.auth.authorization_dependencies import LogTeamLookup, LogTeamLookupDependency from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.list_api.common import ( PROBLEM_TYPE_BASE, @@ -16,7 +19,6 @@ from litellm.proxy.list_api.common import ( reject_unknown_query_params, ) from litellm.proxy.management_endpoints.management_v1.common import MANAGEMENT_V1_PREFIX -from litellm.proxy.utils import PrismaClient from litellm.types.proxy.management_endpoints.management_v1 import ( FacetListResponse, PageMeta, @@ -37,49 +39,29 @@ def _as_utc(value: datetime) -> datetime: async def _spend_log_scope_clause( user_api_key_dict: UserAPIKeyAuth, - prisma_client: PrismaClient, + log_team_lookup: LogTeamLookup, next_param_index: int, -) -> tuple[str | None, tuple[str | list[str], ...]]: +) -> tuple[str | None, tuple[object, ...]]: """SQL predicate restricting the facet to spend logs this caller may read. Returns ``(None, ())`` for a proxy admin. Mirrors the scoping ``/spend/logs/ui`` applies, so a dropdown can never offer a value from a row the caller could not open. """ - from litellm.proxy.spend_tracking.spend_management_endpoints import ( - _get_permitted_team_ids_for_spend_logs, - _is_admin_view_safe, - ) + from litellm.proxy.spend_tracking.spend_management_endpoints import _is_admin_view_safe, read_scope_sql if _is_admin_view_safe(user_api_key_dict=user_api_key_dict): return None, () - - try: - permitted_team_ids = await _get_permitted_team_ids_for_spend_logs( - prisma_client=prisma_client, - user_api_key_dict=user_api_key_dict, - ) - except Exception: - permitted_team_ids = [] - - caller_user_id: Final = user_api_key_dict.user_id - # = ANY(::text[]) rather than an expanded IN list, matching the clause - # ui_view_spend_logs builds: one parameter whatever the team count. - templates: Final = (('"user" = ${}',) if caller_user_id is not None else ()) + ( - ("team_id = ANY(${}::text[])",) if permitted_team_ids else () + scope: Final = await resolve_owned_read_scope( + user_api_key_dict.user_id, partial(log_team_lookup, user_api_key_dict) ) - params: Final = ((caller_user_id,) if caller_user_id is not None else ()) + ( - (permitted_team_ids,) if permitted_team_ids else () - ) - if not templates: - return "FALSE", () - clauses: Final = tuple(template.format(next_param_index + offset) for offset, template in enumerate(templates)) - return f"({' OR '.join(clauses)})", params + return read_scope_sql(scope, next_param_index) async def _list_spend_log_facet( request: Request, user_api_key_dict: UserAPIKeyAuth, + log_team_lookup: LogTeamLookup, start_time: datetime, end_time: datetime, q: str | None, @@ -107,7 +89,7 @@ async def _list_spend_log_facet( scope_clause, scope_params = await _spend_log_scope_clause( user_api_key_dict=user_api_key_dict, - prisma_client=prisma_client, + log_team_lookup=log_team_lookup, next_param_index=len(window_params) + len(search_params) + 1, ) @@ -178,6 +160,7 @@ async def _list_spend_log_facet( async def list_spend_log_end_users( request: Request, user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + log_team_lookup: LogTeamLookupDependency, start_time: Annotated[ datetime, Query(alias="filter[startTime][gte]", description="Window start (UTC when no offset is given)"), @@ -211,6 +194,7 @@ async def list_spend_log_end_users( return await _list_spend_log_facet( request=request, user_api_key_dict=user_api_key_dict, + log_team_lookup=log_team_lookup, start_time=start_time, end_time=end_time, q=q, @@ -229,6 +213,7 @@ async def list_spend_log_end_users( async def list_spend_log_users( request: Request, user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + log_team_lookup: LogTeamLookupDependency, start_time: Annotated[ datetime, Query(alias="filter[startTime][gte]", description="Window start (UTC when no offset is given)"), @@ -245,6 +230,7 @@ async def list_spend_log_users( return await _list_spend_log_facet( request=request, user_api_key_dict=user_api_key_dict, + log_team_lookup=log_team_lookup, start_time=start_time, end_time=end_time, q=q, diff --git a/litellm/proxy/spend_tracking/log_visibility.py b/litellm/proxy/spend_tracking/log_visibility.py deleted file mode 100644 index 83f236d3028..00000000000 --- a/litellm/proxy/spend_tracking/log_visibility.py +++ /dev/null @@ -1,44 +0,0 @@ -from collections.abc import Awaitable, Callable -from dataclasses import dataclass -from typing import Final - -from fastapi import HTTPException - -from litellm.proxy._types import UserAPIKeyAuth - - -@dataclass(frozen=True, slots=True) -class LogVisibility: - all_teams: bool = False - user_id: str = "" - team_ids: tuple[str, ...] = () - api_key_hash: str = "" - - -async def permitted_log_teams(auth: UserAPIKeyAuth) -> tuple[str, ...]: - from litellm.proxy.proxy_server import prisma_client - from litellm.proxy.spend_tracking.spend_management_endpoints import ( - _get_permitted_team_ids_for_spend_logs_or_empty, # pyright: ignore[reportPrivateUsage] # Reuse request-log policy - ) - - if prisma_client is None: - return () - return await _get_permitted_team_ids_for_spend_logs_or_empty(prisma_client=prisma_client, user_api_key_dict=auth) - - -async def log_visibility( - auth: UserAPIKeyAuth, - team_lookup: Callable[[UserAPIKeyAuth], Awaitable[tuple[str, ...]]] = permitted_log_teams, -) -> LogVisibility: - from litellm.proxy.spend_tracking.spend_management_endpoints import ( - _is_admin_view_safe, # pyright: ignore[reportPrivateUsage] # Reuse request-log policy - ) - - if _is_admin_view_safe(user_api_key_dict=auth): - return LogVisibility(all_teams=True) - if auth.user_id: - team_ids: Final = await team_lookup(auth) - return LogVisibility(user_id=auth.user_id, team_ids=team_ids, api_key_hash=auth.token or "") - if auth.token: - return LogVisibility(api_key_hash=auth.token) - raise HTTPException(status_code=403, detail="Not allowed to view logs") diff --git a/litellm/proxy/spend_tracking/spend_management_endpoints.py b/litellm/proxy/spend_tracking/spend_management_endpoints.py index 3102fc63cf4..48cc684549f 100644 --- a/litellm/proxy/spend_tracking/spend_management_endpoints.py +++ b/litellm/proxy/spend_tracking/spend_management_endpoints.py @@ -3,8 +3,8 @@ import collections import json import os from collections.abc import Mapping, Sequence -from dataclasses import dataclass from datetime import date, datetime, timedelta, timezone +from functools import partial from itertools import groupby from types import MappingProxyType from typing import ( @@ -37,6 +37,18 @@ from litellm.constants import ( from litellm.litellm_core_utils.classifier_logging import classifier_audit_fields, classifier_input_snapshot from litellm.proxy._types import * from litellm.proxy._types import ProviderBudgetResponse, ProviderBudgetResponseObject +from litellm.proxy.auth.authorization import ( + AllRows, + OwnedRows, + ReadScope, + can_read_log_owner, + can_read_team_logs, + resolve_owned_read_scope, +) +from litellm.proxy.auth.authorization_dependencies import ( + LogTeamLookup, + LogTeamLookupDependency, +) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.litellm_pre_call_utils import LiteLLMProxyRequestSetup from litellm.proxy.spend_tracking.spend_capture_rate import ( @@ -403,11 +415,6 @@ async def _find_team_row(prisma_client: PrismaClient, team_id: str) -> _Supports return await _team_table(prisma_client).find_unique(where={"team_id": team_id}) -async def _find_team_rows(prisma_client: PrismaClient, team_ids: Sequence[str]) -> Sequence[_SupportsModelDump]: - """Read team rows as Prisma model instances.""" - return await _team_table(prisma_client).find_many(where={"team_id": {"in": team_ids}}) - - @router.get( "/spend/keys", tags=["Budget & Spend Tracking"], @@ -2474,6 +2481,7 @@ def _build_spend_log_search_condition( ) async def ui_view_spend_logs( request: Request, + log_team_lookup: LogTeamLookupDependency, api_key: str | None = fastapi.Query( default=None, description="Get spend logs based on api key", @@ -2775,16 +2783,8 @@ async def ui_view_spend_logs( and team_id is None and (is_request_id_lookup or _can_user_view_spend_log(user_api_key_dict=user_api_key_dict)) ) - permitted_team_ids: Final = ( - await _get_permitted_team_ids_for_spend_logs_or_empty( - prisma_client=prisma_client, - user_api_key_dict=user_api_key_dict, - ) - if user_scope_applies - else () - ) - explicit_user_requires_caller_scope: Final = ( - user_scope_applies and not permitted_team_ids and user_id is not None + read_scope: Final = ( + await _spend_log_read_scope(user_api_key_dict, log_team_lookup) if user_scope_applies else AllRows() ) if not is_admin_view: if team_id is not None: @@ -2799,22 +2799,6 @@ async def ui_view_spend_logs( detail={"error": f"Not authorized to view team spend for team_id={team_id}"}, ) where_conditions["team_id"] = team_id - elif user_scope_applies: - if permitted_team_ids: - if user_id is None: - where_conditions.pop("user", None) - where_conditions["OR"] = [ - {"user": user_api_key_dict.user_id}, - {"team_id": {"in": permitted_team_ids}}, - ] - else: - if user_id is None: - where_conditions["user"] = user_api_key_dict.user_id - else: - where_conditions["AND"] = where_conditions.get("AND", []) + [ - {"user": user_api_key_dict.user_id} - ] - where_conditions.pop("team_id", None) # Calculate skip value for pagination skip: Final = (page - 1) * page_size @@ -2874,17 +2858,11 @@ async def ui_view_spend_logs( sql_params.append(request_id_filter) p += 1 - # Multi-team OR filter: (user = $X OR team_id = ANY($Y)) - if permitted_team_ids: - or_clause: Final = f'("user" = ${p} OR team_id = ANY(${p + 1}::text[]))' - sql_params.append(user_api_key_dict.user_id) - sql_params.append(permitted_team_ids) - p += 2 - sql_conditions.append(or_clause) - elif explicit_user_requires_caller_scope: - sql_conditions.append(f'"user" = ${p}') - sql_params.append(user_api_key_dict.user_id) - p += 1 + scope_clause, scope_params = read_scope_sql(read_scope, p) + if scope_clause: + sql_conditions.append(scope_clause) + sql_params.extend(scope_params) + p += len(scope_params) if session_id is not None and isinstance(session_id, str): like_escaped_session_id: Final = session_id.replace("\\", "\\\\").replace("%", "\\%").replace("_", "\\_") @@ -3390,6 +3368,7 @@ async def _resolve_request_response_payload( ) async def ui_view_request_response_for_request_id( request_id: str, + log_team_lookup: LogTeamLookupDependency, start_date: str | None = fastapi.Query( default=None, description="Time from which to start viewing key spend", @@ -3442,6 +3421,7 @@ async def ui_view_request_response_for_request_id( user_api_key_dict=user_api_key_dict, request_id=request_id, caller_is_admin=caller_is_admin, + log_team_lookup=log_team_lookup, ) ) stored_request_id: Final = _stored_request_id(spend_log_row, request_id) @@ -4512,6 +4492,7 @@ async def ui_get_spend_by_tags( }, ) async def ui_view_session_spend_logs( + log_team_lookup: LogTeamLookupDependency, session_id: str = fastapi.Query( description="Get all spend logs for a particular session", ), @@ -4549,36 +4530,16 @@ async def ui_view_session_spend_logs( detail="Database not connected", ) - if _is_admin_view_safe(user_api_key_dict=user_api_key_dict): - scope_sql = "" - scope_params = () - where_conditions = {"session_id": session_id} - else: - try: - permitted_team_ids = ( - await _get_permitted_team_ids_for_spend_logs( - prisma_client=prisma_client, - user_api_key_dict=user_api_key_dict, - ) - if _can_user_view_spend_log(user_api_key_dict=user_api_key_dict) - else [] - ) - except Exception: # noqa: BLE001 # mirror /spend/logs/ui: failed team lookup falls back to own-logs-only scope - permitted_team_ids = [] - if permitted_team_ids: - scope_sql = ' AND ("user" = $4 OR team_id = ANY($5::text[]))' - scope_params = (user_api_key_dict.user_id, permitted_team_ids) - where_conditions = { - "session_id": session_id, - "OR": [ - {"user": user_api_key_dict.user_id}, - {"team_id": {"in": permitted_team_ids}}, - ], - } - else: - scope_sql = ' AND "user" = $4' - scope_params = (user_api_key_dict.user_id,) - where_conditions = {"session_id": session_id, "user": user_api_key_dict.user_id} + read_scope: Final = ( + AllRows() + if _is_admin_view_safe(user_api_key_dict=user_api_key_dict) + else await _spend_log_read_scope(user_api_key_dict, log_team_lookup) + if _can_user_view_spend_log(user_api_key_dict=user_api_key_dict) + else OwnedRows(user_api_key_dict.user_id) + ) + scope_clause, scope_params = read_scope_sql(read_scope, 4) + scope_sql: Final = f" AND {scope_clause}" if scope_clause else "" + where_conditions: Final = {"session_id": session_id, **_read_scope_where(read_scope)} # Calculate pagination offsets skip: Final = (page - 1) * page_size @@ -4859,22 +4820,12 @@ async def _can_team_member_view_log( Returns True if the team exists and the user is either a team admin or a team member with the ``/spend/logs`` permission. """ - from litellm.proxy.management.teams.access import is_team_admin - from litellm.proxy.management_endpoints.common_utils import _team_member_has_permission - if team_id is None: return False team_row: Final = await _find_team_row(prisma_client, team_id) if team_row is None: return False - team_obj: Final = LiteLLM_TeamTable.model_validate(team_row.model_dump()) - if is_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj): - return True - return _team_member_has_permission( - user_api_key_dict=user_api_key_dict, - team_obj=team_obj, - permission=KeyManagementRoutes.SPEND_LOGS.value, - ) + return can_read_team_logs(user_api_key_dict, LiteLLM_TeamTable.model_validate(team_row.model_dump())) def _can_user_view_spend_log(user_api_key_dict: UserAPIKeyAuth) -> bool: @@ -4899,15 +4850,12 @@ async def _user_can_view_spend_log_owner( owner_user: str | None, owner_team_id: str | None, ) -> bool: - if owner_user is not None and owner_user == user_api_key_dict.user_id: - return True - if owner_team_id: - return await _can_team_member_view_log( - prisma_client=prisma_client, - user_api_key_dict=user_api_key_dict, - team_id=owner_team_id, - ) - return False + return await can_read_log_owner( + user_api_key_dict.user_id, + owner_user, + owner_team_id, + partial(_can_team_member_view_log, prisma_client, user_api_key_dict), + ) def _spend_log_forbidden(request_id: str) -> HTTPException: @@ -4940,44 +4888,50 @@ async def _assert_user_can_view_request_id( raise _spend_log_forbidden(request_id) -@dataclass(frozen=True, slots=True) -class _SpendLogViewer: - user_id: str | None - team_ids: tuple[str, ...] - - -async def _spend_log_viewer(prisma_client: PrismaClient, user_api_key_dict: UserAPIKeyAuth) -> _SpendLogViewer: - return _SpendLogViewer( - user_id=user_api_key_dict.user_id, - team_ids=await _get_permitted_team_ids_for_spend_logs_or_empty( - prisma_client=prisma_client, - user_api_key_dict=user_api_key_dict, - ), +async def _spend_log_read_scope(user_api_key_dict: UserAPIKeyAuth, log_team_lookup: LogTeamLookup) -> OwnedRows: + return await resolve_owned_read_scope( + user_api_key_dict.user_id, + partial(log_team_lookup, user_api_key_dict), ) -def _viewer_scope_clause(viewer: _SpendLogViewer | None) -> tuple[str, tuple[object, ...]]: - match viewer: - case None: - return ("", ()) - case _SpendLogViewer(user_id=user_id, team_ids=()): - return (' AND "user" = $2', (user_id,)) - case _SpendLogViewer(user_id=user_id, team_ids=team_ids): - return (' AND ("user" = $2 OR team_id = ANY($3::text[]))', (user_id, team_ids)) +def read_scope_sql(scope: ReadScope, next_param: int) -> tuple[str, tuple[object, ...]]: + if isinstance(scope, AllRows): + return "", () + if scope.user_id is not None and scope.team_ids: + return ( + f'("user" = ${next_param} OR team_id = ANY(${next_param + 1}::text[]))', + (scope.user_id, scope.team_ids), + ) + if scope.user_id is not None: + return f'"user" = ${next_param}', (scope.user_id,) + if scope.team_ids: + return f"team_id = ANY(${next_param}::text[])", (scope.team_ids,) + return "FALSE", () -def _spend_log_payload_query(request_id: str, viewer: _SpendLogViewer | None) -> tuple[str, tuple[object, ...]]: +def _read_scope_where(scope: ReadScope) -> Mapping[str, object]: + if isinstance(scope, AllRows): + return {} + user_grant: Final = ({"user": scope.user_id},) if scope.user_id is not None else () + team_grant: Final = ({"team_id": {"in": list(scope.team_ids)}},) if scope.team_ids else () + grants: Final = user_grant + team_grant + return grants[0] if len(grants) == 1 else {"OR": list(grants)} + + +def _spend_log_payload_query(request_id: str, scope: ReadScope) -> tuple[str, tuple[object, ...]]: """ Fetch the one row an id lookup resolves to, preferring the exact ``request_id`` match over rows that merely carry the id as their client-set ``litellm_call_id``. A non-admin viewer only ever gets rows they own or rows of a team they may view. """ - scope, scope_params = _viewer_scope_clause(viewer) + scope_clause, scope_params = read_scope_sql(scope, 2) + scope_sql: Final = f" AND {scope_clause}" if scope_clause else "" return ( f""" SELECT request_id, messages, response, proxy_server_request, metadata, "user", team_id FROM "LiteLLM_SpendLogs" - WHERE (request_id = $1 OR litellm_call_id = $1){scope} + WHERE (request_id = $1 OR litellm_call_id = $1){scope_sql} ORDER BY (request_id = $1) DESC LIMIT 1 """, @@ -4990,6 +4944,7 @@ async def _resolve_spend_log_payload_row( user_api_key_dict: UserAPIKeyAuth, request_id: str, caller_is_admin: bool, + log_team_lookup: LogTeamLookup, ) -> Mapping[str, object] | None: """ Resolve an id lookup to the caller's own spend-log row before any payload @@ -4998,8 +4953,8 @@ async def _resolve_spend_log_payload_row( that id is only the caller's ``litellm_call_id``; the row's stored ``request_id`` is the key that names the caller's own request. """ - viewer: Final = None if caller_is_admin else await _spend_log_viewer(prisma_client, user_api_key_dict) - sql_query, sql_params = _spend_log_payload_query(request_id, viewer) + scope: Final = AllRows() if caller_is_admin else await _spend_log_read_scope(user_api_key_dict, log_team_lookup) + sql_query, sql_params = _spend_log_payload_query(request_id, scope) rows: Final[Sequence[Mapping[str, object]] | None] = await _query_raw_or_none(prisma_client, sql_query, *sql_params) if not rows: return None @@ -5075,57 +5030,3 @@ async def _assert_user_owns_cold_storage_payload( owner_user, owner_team_id = _cold_storage_payload_owner(payload) if not await _user_can_view_spend_log_owner(prisma_client, user_api_key_dict, owner_user, owner_team_id): raise _spend_log_forbidden(request_id) - - -async def _get_permitted_team_ids_for_spend_logs( - prisma_client: PrismaClient, - user_api_key_dict: UserAPIKeyAuth, -) -> list[str]: - """ - Return team IDs where the user is either a team admin or has the - ``/spend/logs`` permission, allowing them to view team-wide spend logs. - """ - # Imported here to avoid circular import: proxy_server imports this module. - from litellm.proxy.auth.auth_checks import get_user_object - from litellm.proxy.management.teams.access import is_team_admin - from litellm.proxy.management_endpoints.common_utils import _team_member_has_permission - from litellm.proxy.proxy_server import proxy_logging_obj, user_api_key_cache - - user_obj: Final = await get_user_object( - user_id=user_api_key_dict.user_id, - prisma_client=prisma_client, - user_api_key_cache=user_api_key_cache, - user_id_upsert=False, - proxy_logging_obj=proxy_logging_obj, - ) - if user_obj is None or not user_obj.teams: - return [] - - team_rows: Final = await _find_team_rows(prisma_client, user_obj.teams) - - permitted: Final[list[str]] = [] - for team_row in team_rows: - team_obj = LiteLLM_TeamTable.model_validate(team_row.model_dump()) - if is_team_admin(user_api_key_dict=user_api_key_dict, team_obj=team_obj) or _team_member_has_permission( - user_api_key_dict=user_api_key_dict, - team_obj=team_obj, - permission=KeyManagementRoutes.SPEND_LOGS.value, - ): - permitted.append(team_obj.team_id) - return permitted - - -async def _get_permitted_team_ids_for_spend_logs_or_empty( - prisma_client: PrismaClient, - user_api_key_dict: UserAPIKeyAuth, -) -> tuple[str, ...]: - """Resolve permitted teams once, falling back to the caller's own-user scope.""" - try: - return tuple( - await _get_permitted_team_ids_for_spend_logs( - prisma_client=prisma_client, - user_api_key_dict=user_api_key_dict, - ) - ) - except Exception: - return () diff --git a/litellm/proxy/tracing_endpoints.py b/litellm/proxy/tracing_endpoints.py index d741e0b29b4..046bcc9f704 100644 --- a/litellm/proxy/tracing_endpoints.py +++ b/litellm/proxy/tracing_endpoints.py @@ -10,6 +10,7 @@ GET /v1/traces/{trace_id}/spans/{span_id} SpanDetail import time from collections.abc import Mapping from dataclasses import dataclass +from functools import partial from http.client import responses from types import MappingProxyType from typing import Annotated, Final @@ -20,12 +21,13 @@ from pydantic import BaseModel, ConfigDict from litellm._logging import verbose_proxy_logger from litellm.constants import OTLP_RETRY_AFTER_SECONDS from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth +from litellm.proxy.auth.authorization import AllRows, ReadScope, resolve_trace_read_scope +from litellm.proxy.auth.authorization_dependencies import LogTeamLookupDependency from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.http_parsing_utils import is_otlp_trace_request -from litellm.proxy.spend_tracking.log_visibility import log_visibility from litellm.proxy.tracing_runtime import provide_receiver, require_receiver from litellm.rust_bridge.trace_query_responses import TraceQueryHelp, TraceSQLResponse -from litellm.rust_bridge.traces import ClickHouseStorage, QueryScope +from litellm.rust_bridge.traces import AllQueryScope, ClickHouseStorage, OwnedQueryScope, QueryScope from litellm.tracing import ( Tenant, TraceReceiver, @@ -43,14 +45,14 @@ MS_PER_DAY: Final = 24 * 60 * 60 * 1000 @dataclass(frozen=True, slots=True) class TraceAccessContext: receiver: TraceReceiver | None - read_scope: TraceScope | None + read_scope: ReadScope | None write_tenant: Tenant | None def reader(self) -> tuple[TraceReceiver, TraceScope]: tracing: Final = require_receiver(self.receiver) if self.read_scope is None: raise HTTPException(status_code=403, detail="Not allowed to view agent traces") - return tracing, self.read_scope + return tracing, _trace_scope(self.read_scope) def writer(self) -> tuple[TraceReceiver, Tenant]: if self.write_tenant is None: @@ -61,27 +63,23 @@ class TraceAccessContext: async def provide_trace_access( auth: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], tracing: Annotated[TraceReceiver | None, Depends(provide_receiver)], + log_team_lookup: LogTeamLookupDependency, ) -> TraceAccessContext: tenant: Final = Tenant( team_id=auth.team_id or "", api_key_hash=auth.token or "", org_id=auth.org_id or "", user_id=auth.user_id or "" ) write_tenant: Final = None if auth.user_role == LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY else tenant - if ( - not auth.user_id - and not auth.token - and auth.user_role not in (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) - ): - return TraceAccessContext(tracing, None, write_tenant) - visibility: Final = await log_visibility(auth) - return TraceAccessContext( - tracing, - TraceScope( - all_teams=1 if visibility.all_teams else 0, - user_id=visibility.user_id, - team_ids=visibility.team_ids, - api_key_hash=visibility.api_key_hash, - ), - write_tenant, + read_scope: Final = await resolve_trace_read_scope(auth, partial(log_team_lookup, auth)) + return TraceAccessContext(tracing, read_scope, write_tenant) + + +def _trace_scope(scope: ReadScope) -> TraceScope: + if isinstance(scope, AllRows): + return TraceScope(all_teams=1, user_id="", team_ids=()) + return TraceScope( + all_teams=0, + user_id=scope.user_id or "", + team_ids=scope.team_ids, ) @@ -160,7 +158,7 @@ class TraceQueryRequest(BaseModel): @dataclass(frozen=True, slots=True) class TraceQueryAccess: storage: ClickHouseStorage - scope: QueryScope + scope: ReadScope secret: str @@ -172,24 +170,27 @@ def provide_trace_query_secret() -> str: return master_key -async def trace_query_scope(auth: UserAPIKeyAuth) -> QueryScope: - visibility: Final = await log_visibility(auth) - if visibility.all_teams: - return {"kind": "admin"} - return { - "kind": "logs", - "user_id": visibility.user_id, - "team_ids": visibility.team_ids, - "api_key_hash": visibility.api_key_hash, - } +def trace_query_scope(scope: ReadScope) -> QueryScope: + if isinstance(scope, AllRows): + return AllQueryScope(kind="all") + return OwnedQueryScope( + kind="owned", + user_id=scope.user_id or "", + team_ids=scope.team_ids, + ) async def provide_trace_query_access( auth: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], tracing: Annotated[TraceReceiver | None, Depends(provide_receiver)], secret: Annotated[str, Depends(provide_trace_query_secret)], + log_team_lookup: LogTeamLookupDependency, ) -> TraceQueryAccess: - return TraceQueryAccess(require_receiver(tracing).store.storage, await trace_query_scope(auth), secret) + storage: Final = require_receiver(tracing).store.storage + scope: Final = await resolve_trace_read_scope(auth, partial(log_team_lookup, auth)) + if scope is None: + raise HTTPException(status_code=403, detail="Not allowed to view logs") + return TraceQueryAccess(storage, scope, secret) @router.post("/v1/traces/query", response_model=TraceSQLResponse, response_model_exclude_unset=True) @@ -198,7 +199,7 @@ async def query_agent_traces( access: Annotated[TraceQueryAccess, Depends(provide_trace_query_access)], ) -> TraceSQLResponse: try: - return await access.storage.query_sql(body.sql, access.scope, access.secret) + return await access.storage.query_sql(body.sql, trace_query_scope(access.scope), access.secret) except ValueError as error: raise HTTPException(status_code=400, detail=str(error)) from error except RuntimeError as error: @@ -211,7 +212,7 @@ async def help_agent_trace_queries( access: Annotated[TraceQueryAccess, Depends(provide_trace_query_access)], ) -> TraceQueryHelp: try: - return await access.storage.query_help(access.scope, access.secret) + return await access.storage.query_help(trace_query_scope(access.scope), access.secret) except RuntimeError as error: verbose_proxy_logger.warning("Trace query help unavailable: %s", error) raise HTTPException(status_code=503, detail="Trace query help is temporarily unavailable") from error diff --git a/litellm/rust_bridge/trace_queries.py b/litellm/rust_bridge/trace_queries.py index 0a9fe186646..1add1787241 100644 --- a/litellm/rust_bridge/trace_queries.py +++ b/litellm/rust_bridge/trace_queries.py @@ -30,7 +30,6 @@ class ListTracesParams(BaseModel): all_teams: Literal[0, 1] user_id: str team_ids: tuple[str, ...] - api_key_hash: str start_ms: Int64 end_ms: Int64 cursor_ms: Int64 @@ -43,7 +42,6 @@ class TraceSpansParams(BaseModel): all_teams: Literal[0, 1] user_id: str team_ids: tuple[str, ...] - api_key_hash: str trace_id: str trace_ref: str @@ -53,7 +51,6 @@ class SpanDetailParams(BaseModel): all_teams: Literal[0, 1] user_id: str team_ids: tuple[str, ...] - api_key_hash: str trace_id: str trace_ref: str span_id: str @@ -64,7 +61,6 @@ class SpanErrorParams(BaseModel): all_teams: Literal[0, 1] user_id: str team_ids: tuple[str, ...] - api_key_hash: str trace_id: str trace_ref: str span_id: str @@ -77,7 +73,6 @@ class SpendByResponseIdsParams(BaseModel): all_teams: Literal[0, 1] user_id: str team_ids: tuple[str, ...] - api_key_hash: str response_ids: tuple[str, ...] start_ms: Int64 end_ms: Int64 @@ -143,7 +138,6 @@ class TraceIdentityParams(BaseModel): all_teams: Literal[0, 1] user_id: str team_ids: tuple[str, ...] - api_key_hash: str trace_id: str diff --git a/litellm/rust_bridge/traces.py b/litellm/rust_bridge/traces.py index 23778983d75..f70b5ec3f93 100644 --- a/litellm/rust_bridge/traces.py +++ b/litellm/rust_bridge/traces.py @@ -77,29 +77,17 @@ class DecodedSpan(TypedDict): consumed_attributes: ReadOnly[tuple[str, str]] -class AdminQueryScope(TypedDict): - kind: ReadOnly[Literal["admin"]] +class AllQueryScope(TypedDict): + kind: ReadOnly[Literal["all"]] -class TeamQueryScope(TypedDict): - kind: ReadOnly[Literal["team"]] - team_id: ReadOnly[str] - - -class KeyQueryScope(TypedDict): - kind: ReadOnly[Literal["key"]] - team_id: ReadOnly[str] - api_key_hash: ReadOnly[str] - - -class LogQueryScope(TypedDict): - kind: ReadOnly[Literal["logs"]] +class OwnedQueryScope(TypedDict): + kind: ReadOnly[Literal["owned"]] user_id: ReadOnly[str] team_ids: ReadOnly[tuple[str, ...]] - api_key_hash: ReadOnly[str] -QueryScope = AdminQueryScope | TeamQueryScope | KeyQueryScope | LogQueryScope +QueryScope = AllQueryScope | OwnedQueryScope class NativeStore(Protocol): diff --git a/litellm/tracing/types.py b/litellm/tracing/types.py index 15b28293590..e177094d464 100644 --- a/litellm/tracing/types.py +++ b/litellm/tracing/types.py @@ -108,7 +108,6 @@ class TraceScope(TypedDict): all_teams: ReadOnly[Literal[0, 1]] user_id: ReadOnly[str] team_ids: ReadOnly[tuple[str, ...]] - api_key_hash: ReadOnly[str] class SpanRow(TypedDict): diff --git a/tests/code_coverage_tests/unbounded_in_baseline.txt b/tests/code_coverage_tests/unbounded_in_baseline.txt index c42a6b0ddf5..01d8760855f 100644 --- a/tests/code_coverage_tests/unbounded_in_baseline.txt +++ b/tests/code_coverage_tests/unbounded_in_baseline.txt @@ -125,9 +125,8 @@ litellm/proxy/proxy_server.py _fetch_db_models_for_search prisma not.in `list(db litellm/proxy/proxy_server.py _gather_team_accessible_model_ids prisma model_name.in `_resolved_names` 0 litellm/proxy/proxy_server.py get_all_team_models prisma team_id.in `user_teams` 0 litellm/proxy/spend_tracking/ptu_flat_cost_rollup.py _prune_filter prisma model.in `chunk` 0 -litellm/proxy/spend_tracking/spend_management_endpoints.py _find_team_rows prisma team_id.in `team_ids` 0 -litellm/proxy/spend_tracking/spend_management_endpoints.py ui_view_session_spend_logs prisma team_id.in `permitted_team_ids` 0 -litellm/proxy/spend_tracking/spend_management_endpoints.py ui_view_spend_logs prisma team_id.in `permitted_team_ids` 0 +litellm/proxy/auth/authorization_dependencies.py load_permitted_log_team_ids prisma team_id.in `user_obj.teams` 0 +litellm/proxy/spend_tracking/spend_management_endpoints.py _read_scope_where prisma team_id.in `list(scope.team_ids)` 0 litellm/proxy/ui_crud_endpoints/proxy_setting_endpoints.py _validate_default_teams_exist prisma team_id.in `list(team_ids)` 0 litellm/proxy/utils.py PrismaClient.check_view_exists raw-sql viewname.IN `IN ( {expected_views_str} )` 0 litellm/proxy/utils.py PrismaClient.delete_data prisma team_id.in `team_id_list` 0 diff --git a/tests/integration/spend/test_spend_log_read_scope.py b/tests/integration/spend/test_spend_log_read_scope.py new file mode 100644 index 00000000000..f9034371e35 --- /dev/null +++ b/tests/integration/spend/test_spend_log_read_scope.py @@ -0,0 +1,226 @@ +import os +import uuid +from collections.abc import AsyncIterator +from dataclasses import dataclass +from datetime import datetime, timedelta, timezone +from typing import Final +from urllib.parse import parse_qsl, urlencode, urlsplit, urlunsplit + +import psycopg +import pytest +import pytest_asyncio +from integration._support.client import Gateway +from prisma import Prisma +from psycopg import sql +from psycopg.types.json import Jsonb +from pydantic import TypeAdapter + +from litellm.proxy.auth.authorization import AllRows, OwnedRows, ReadScope +from litellm.proxy.spend_tracking.spend_management_endpoints import _spend_log_payload_query, read_scope_sql + + +@dataclass(frozen=True, slots=True) +class SpendRow: + request_id: str + user: str | None + team_id: str | None + call_id: str | None = None + + +@dataclass(frozen=True, slots=True) +class RequestId: + request_id: str + + +REQUEST_IDS: Final = TypeAdapter(tuple[RequestId, ...]) +ROWS: Final = ( + SpendRow("own", "caller", None, "foreign"), + SpendRow("team-1", "other", "first"), + SpendRow("team-2", "third", "second"), + SpendRow("foreign", "other", "outside"), + SpendRow("ownerless", None, None), + SpendRow("team-ownerless", None, "first"), +) + + +def _seed_rows( + connection: psycopg.Connection, + schema: str, + rows: tuple[SpendRow, ...], + session_id: str, + started: datetime, +) -> None: + utc_timestamp: Final = started.astimezone(timezone.utc).replace(tzinfo=None) + with connection.cursor() as cursor: + cursor.executemany( + sql.SQL( + 'INSERT INTO {} (request_id, "user", team_id, litellm_call_id, session_id, ' + '"startTime", "endTime", messages, response, call_type) ' + "VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, 'acompletion')" + ).format(sql.Identifier(schema, "LiteLLM_SpendLogs")), + tuple( + ( + row.request_id, + row.user, + row.team_id, + row.call_id, + session_id, + utc_timestamp, + utc_timestamp, + Jsonb([{"role": "user", "content": row.request_id + " payload"}]), + Jsonb({"id": row.request_id}), + ) + for row in rows + ), + ) + + +@pytest_asyncio.fixture(loop_scope="function") +async def spend_database() -> AsyncIterator[Prisma]: + schema: Final = f"integration_spend_scope_{uuid.uuid4().hex}" + url: Final = os.environ["DATABASE_URL"] + parsed: Final = urlsplit(url) + scoped_url: Final = urlunsplit( + parsed._replace(query=urlencode({**dict(parse_qsl(parsed.query)), "schema": schema})) + ) + with psycopg.connect(url, autocommit=True) as setup: + setup.execute(sql.SQL("CREATE SCHEMA {}").format(sql.Identifier(schema))) + try: + setup.execute( + sql.SQL('CREATE TABLE {} (LIKE public."LiteLLM_SpendLogs" INCLUDING ALL)').format( + sql.Identifier(schema, "LiteLLM_SpendLogs") + ) + ) + _seed_rows(setup, schema, ROWS, "scope-session", datetime(2026, 1, 1, tzinfo=timezone.utc)) + database: Final = Prisma(datasource={"url": scoped_url}) + await database.connect() + try: + yield database + finally: + await database.disconnect() + finally: + setup.execute(sql.SQL("DROP SCHEMA {} CASCADE").format(sql.Identifier(schema))) + + +@pytest.mark.asyncio +@pytest.mark.parametrize("preceding_filters", [False, True]) +@pytest.mark.parametrize( + ("scope", "user_filter", "expected"), + [ + (AllRows(), None, ("foreign", "own", "ownerless", "team-1", "team-2", "team-ownerless")), + (OwnedRows("caller"), None, ("own",)), + (OwnedRows(None), None, ()), + (OwnedRows(None, ("first", "second")), None, ("team-1", "team-2", "team-ownerless")), + (OwnedRows(None, ("first", "second")), "other", ("team-1",)), + (OwnedRows("caller", ("first", "second")), None, ("own", "team-1", "team-2", "team-ownerless")), + (OwnedRows("caller", ("first", "second")), "other", ("team-1",)), + (OwnedRows("caller", ("first' OR TRUE --",)), None, ("own",)), + (OwnedRows("caller' OR TRUE --", ("first",)), None, ("team-1", "team-ownerless")), + ], +) +async def test_ownership_sql_selects_allowed_rows_and_intersects_filters( + spend_database: Prisma, + scope: ReadScope, + user_filter: str | None, + expected: tuple[str, ...], + preceding_filters: bool, +) -> None: + window_params: Final = ("scope-session", "2026-01-01", "2026-01-02") if preceding_filters else () + window_sql: Final = ( + 'session_id = $1 AND "startTime" >= $2::timestamp AND "startTime" < $3::timestamp AND ' + if preceding_filters + else "" + ) + clause, scope_params = read_scope_sql(scope, len(window_params) + 1) + filter_sql: Final = f' AND "user" = ${len(window_params) + len(scope_params) + 1}' if user_filter else "" + params: Final = window_params + scope_params + ((user_filter,) if user_filter else ()) + result: Final = await spend_database.query_raw( + f'SELECT request_id FROM "LiteLLM_SpendLogs" WHERE {window_sql}{clause or "TRUE"}{filter_sql} ' + "ORDER BY request_id", + *params, + ) + assert tuple(row.request_id for row in REQUEST_IDS.validate_python(result)) == expected + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("scope", "expected"), + [(AllRows(), ("foreign",)), (OwnedRows("caller"), ("own",)), (OwnedRows(None), ())], +) +async def test_payload_sql_filters_foreign_collisions_and_prefers_exact_ids_for_admins( + spend_database: Prisma, scope: ReadScope, expected: tuple[str, ...] +) -> None: + query, params = _spend_log_payload_query("foreign", scope) + result: Final = await spend_database.query_raw(query, *params) + assert tuple(row.request_id for row in REQUEST_IDS.validate_python(result)) == expected + + +def _delete_session(session_id: str) -> None: + with psycopg.connect(os.environ["DATABASE_URL"]) as connection: + connection.execute('DELETE FROM "LiteLLM_SpendLogs" WHERE session_id = %s', (session_id,)) + + +@pytest.mark.parametrize( + ("member_role", "permissions", "team_access"), + [ + ("admin", [], True), + ("user", ["/spend/logs"], True), + ("user", ["/key/info"], False), + ("user", [], False), + ], +) +def test_spend_log_routes_preserve_user_and_permitted_team_access( + gateway: Gateway, member_role: str, permissions: list[str], team_access: bool +) -> None: + session_id: Final = f"scope-{uuid.uuid4().hex}" + started: Final = datetime.now(timezone.utc) - timedelta(hours=1) + with gateway.scenario() as scenario: + caller: Final = scenario.user(user_role="internal_user") + other: Final = scenario.user(user_role="internal_user") + team: Final = scenario.team( + members_with_roles=[{"user_id": caller, "role": member_role}], + team_member_permissions=list(permissions), + ) + outside_team: Final = scenario.team( + members_with_roles=[{"user_id": other, "role": "admin"}], + team_member_permissions=["/spend/logs"], + ) + key: Final = scenario.key(user_id=caller) + other_key: Final = scenario.key(user_id=other) + rows: Final = ( + SpendRow(session_id + "-own", caller, None, session_id + "-foreign"), + SpendRow(session_id + "-team", other, team), + SpendRow(session_id + "-foreign", other, outside_team), + SpendRow(session_id + "-ownerless", None, None), + SpendRow(session_id + "-outside", other, outside_team), + ) + scenario.cleanups.callback(_delete_session, session_id) + with psycopg.connect(os.environ["DATABASE_URL"]) as connection: + _seed_rows(connection, "public", rows, session_id, started) + expected: Final = (rows[0].request_id, rows[1].request_id) if team_access else (rows[0].request_id,) + session: Final = gateway.request("GET", "/spend/logs/session/ui", key=key, params={"session_id": session_id}) + assert session.status_code == 200, session.text + assert session.json()["total"] == len(expected), session.text + assert sorted(row["request_id"] for row in session.json()["data"]) == list(expected), session.text + filters: Final = { + "session_id": session_id, + "start_date": (started - timedelta(hours=1)).strftime("%Y-%m-%d %H:%M:%S"), + "end_date": (started + timedelta(hours=1)).strftime("%Y-%m-%d %H:%M:%S"), + } + listed: Final = gateway.request("GET", "/spend/logs/ui", key=key, params=filters) + assert listed.status_code == 200, listed.text + assert sorted(row["request_id"] for row in listed.json()["data"]) == list(expected), listed.text + narrowed: Final = gateway.request("GET", "/spend/logs/ui", key=key, params={**filters, "user_id": other}) + assert narrowed.status_code == 200, narrowed.text + assert [row["request_id"] for row in narrowed.json()["data"]] == ( + [rows[1].request_id] if team_access else [] + ), narrowed.text + refused: Final = gateway.request("GET", f"/spend/logs/ui/{rows[4].request_id}", key=key) + assert refused.status_code == 403, refused.text + for caller_key, expected_id in ((key, rows[0].request_id), (other_key, rows[2].request_id)): + payload: Final = gateway.request("GET", f"/spend/logs/ui/{rows[2].request_id}", key=caller_key) + assert payload.status_code == 200, payload.text + assert payload.json()["messages"] == [{"role": "user", "content": expected_id + " payload"}], payload.text + admin: Final = gateway.request("GET", f"/spend/logs/ui/{rows[2].request_id}") + assert admin.status_code == 200, admin.text + assert admin.json()["messages"] == [{"role": "user", "content": rows[2].request_id + " payload"}], admin.text diff --git a/tests/test_litellm/tracing/test_receiver.py b/tests/test_litellm/tracing/test_receiver.py index b8a92606417..36a10ea1e42 100644 --- a/tests/test_litellm/tracing/test_receiver.py +++ b/tests/test_litellm/tracing/test_receiver.py @@ -106,7 +106,7 @@ async def test_empty_export_writes_nothing(): async def test_reads_delegate_to_store(): store = _fake_store() tracing = TraceReceiver(store) - scope: TraceScope = {"team_ids": ("team-research",), "api_key_hash": ""} + scope: Final[TraceScope] = {"all_teams": 0, "user_id": "", "team_ids": ("team-research",)} assert await tracing.get_trace("t1", scope) is None store.get_trace.assert_awaited_once_with("t1", scope, "") diff --git a/tests/test_litellm/tracing/test_store.py b/tests/test_litellm/tracing/test_store.py index 0c1820d7c5f..18a5865db7f 100644 --- a/tests/test_litellm/tracing/test_store.py +++ b/tests/test_litellm/tracing/test_store.py @@ -354,7 +354,7 @@ async def test_list_traces_sets_next_cursor_on_full_page(): } client.query = AsyncMock(return_value=[row, {**row, "trace_id": "t1", "trace_ref": "ref1", "start_ms": 900}]) store = TraceStore(client) - scope: TraceScope = {"all_teams": 0, "user_id": "", "team_ids": ("team-a",), "api_key_hash": ""} + scope: Final[TraceScope] = {"all_teams": 0, "user_id": "", "team_ids": ("team-a",)} page = await store.list_traces(scope, 0, 2000, limit=2) assert [t["trace_id"] for t in page["data"]] == ["t2", "t1"] @@ -373,7 +373,7 @@ async def test_get_span_not_found_and_found(): client = MagicMock() client.query = AsyncMock(return_value=[]) store = TraceStore(client) - scope: TraceScope = {"all_teams": 1, "user_id": "", "team_ids": (), "api_key_hash": ""} + scope: Final[TraceScope] = {"all_teams": 1, "user_id": "", "team_ids": ()} assert await store.get_span("t", "s", scope, "ref") is None stored_input = '[{"role": "user", "content": "hi"}]' client.query = AsyncMock( @@ -425,7 +425,7 @@ async def test_trace_cost_is_scoped_and_counts_repeated_request_once(): ] client.query = AsyncMock(side_effect=[spans, tuple(SpendRow.model_validate({**row, "user": ""}) for row in spend)]) store = TraceStore(client) - scope: TraceScope = {"all_teams": 0, "user_id": "", "team_ids": ("team-a",), "api_key_hash": ""} + scope: Final[TraceScope] = {"all_teams": 0, "user_id": "", "team_ids": ("team-a",)} trace = await store.get_trace("trace-1", scope, "ref") @@ -473,7 +473,7 @@ async def test_run_list_uses_matching_spend_and_leaves_missing_cost_unavailable( } ] client.query = AsyncMock(side_effect=[rows, tuple(SpendRow.model_validate({**row, "user": ""}) for row in spend)]) - scope: TraceScope = {"all_teams": 0, "user_id": "", "team_ids": ("team-a",), "api_key_hash": ""} + scope: Final[TraceScope] = {"all_teams": 0, "user_id": "", "team_ids": ("team-a",)} page = await TraceStore(client).list_traces(scope, 0, 2000) @@ -484,7 +484,7 @@ async def test_run_list_uses_matching_spend_and_leaves_missing_cost_unavailable( @pytest.mark.asyncio async def test_ambiguous_cache_response_id_keeps_cost_unavailable(): client = MagicMock() - span = _llm_row("llm-1", "", "agent", "response-1", team_id="", api_key_hash="key-a") + span: Final = _llm_row("llm-1", "", "agent", "response-1", team_id="", user_id="user", api_key_hash="key-a") spend = [ { "request_id": request_id, @@ -496,9 +496,11 @@ async def test_ambiguous_cache_response_id_keeps_cost_unavailable(): } for request_id, cost in (("response-1", 0.25), ("response-1_cache_hit123", 0.0)) ] - client.query = AsyncMock(side_effect=[[span], tuple(SpendRow.model_validate({**row, "user": ""}) for row in spend)]) + client.query = AsyncMock( + side_effect=[[span], tuple(SpendRow.model_validate({**row, "user": "user"}) for row in spend)] + ) store = TraceStore(client) - scope: TraceScope = {"all_teams": 0, "user_id": "", "team_ids": (), "api_key_hash": "key-a"} + scope: Final[TraceScope] = {"all_teams": 0, "user_id": "user", "team_ids": ()} trace = await store.get_trace("trace-1", scope, "ref") @@ -521,7 +523,7 @@ async def test_diagnostic_continuation_preserves_content_version_scope_and_unico ] ) store = TraceStore(client) - scope = {"all_teams": 0, "user_id": "", "team_ids": ("team-a",), "api_key_hash": "key-a"} + scope: Final[TraceScope] = {"all_teams": 0, "user_id": "", "team_ids": ("team-a",)} first = await store.get_span_error("trace-1", "span-1", scope, "scoped-run") assert first is not None and first["next_cursor"] is not None last = await store.get_span_error("trace-1", "span-1", scope, "scoped-run", first["next_cursor"]) @@ -548,7 +550,7 @@ async def test_malformed_diagnostic_cursor_never_reaches_storage(cursor): client.query = AsyncMock() with pytest.raises(ValueError, match="Invalid diagnostic cursor"): await TraceStore(client).get_span_error( - "trace", "span", {"all_teams": 1, "user_id": "", "team_ids": (), "api_key_hash": ""}, cursor=cursor + "trace", "span", {"all_teams": 1, "user_id": "", "team_ids": ()}, cursor=cursor ) client.query.assert_not_awaited() @@ -660,7 +662,7 @@ async def test_trace_id_collision_requires_a_visible_reference_before_reading_co storage: Final = MagicMock() storage.query = AsyncMock(return_value=(TraceIdentityRow(trace_ref="first"), TraceIdentityRow(trace_ref="second"))) store: Final = TraceStore(storage) - scope: Final[TraceScope] = {"all_teams": 1, "user_id": "", "team_ids": (), "api_key_hash": ""} + scope: Final[TraceScope] = {"all_teams": 1, "user_id": "", "team_ids": ()} with pytest.raises(AmbiguousTraceError, match="provide trace_ref"): await store.get_trace("shared-id", scope) with pytest.raises(AmbiguousTraceError, match="provide trace_ref"): diff --git a/tests/test_litellm_rust/test_traces.py b/tests/test_litellm_rust/test_traces.py index 264f489520d..e0208312f18 100644 --- a/tests/test_litellm_rust/test_traces.py +++ b/tests/test_litellm_rust/test_traces.py @@ -43,6 +43,7 @@ def span_row() -> dict[str, JsonValue]: "name": "root", "type": "agent", "agent": "", + "framework": "", "status": "STATUS_CODE_OK", "status_message": "", "error_truncated": 0, @@ -62,7 +63,7 @@ def span_row() -> dict[str, JsonValue]: @pytest.fixture def span_params() -> dict[str, str | int | list[str]]: - return {"trace_id": "trace-1", "trace_ref": "", "all_teams": 1, "user_id": "", "team_ids": [], "api_key_hash": ""} + return {"trace_id": "trace-1", "trace_ref": "", "all_teams": 1, "user_id": "", "team_ids": []} @pytest.mark.asyncio @@ -128,7 +129,7 @@ async def test_from_env_reads_with_clickhouse_url( recording_server.enqueue(ResponseSpec(body={"data": []})) monkeypatch.setenv("CLICKHOUSE_URL", recording_server.base_url) monkeypatch.delenv("CLICKHOUSE_READER_URL", raising=False) - scope: Final[TraceScope] = {"all_teams": 1, "user_id": "", "team_ids": (), "api_key_hash": ""} + scope: Final[TraceScope] = {"all_teams": 1, "user_id": "", "team_ids": ()} page: Final = await TraceReceiver.from_env().list_traces(scope, 0, 1) assert page == {"data": (), "next_cursor": None} assert len(recording_server.requests) == 1 @@ -295,14 +296,23 @@ async def test_insert_validates_values_without_pydantic_copy(recording_server: R assert stored["SpanAttributes"] == attributes -@pytest.mark.parametrize("role", ["proxy_admin", "proxy_admin_viewer", "internal_user"]) -def test_trace_sql_endpoint_executes_for_admin_and_preserves_clickhouse_envelope( - recording_server: RecordingServer, role: str +@pytest.mark.parametrize( + ("role", "user_id", "expected_status"), + ( + ("proxy_admin", None, 200), + ("proxy_admin_viewer", None, 200), + ("internal_user", "user", 200), + ("internal_user", None, 403), + ), +) +def test_trace_sql_endpoint_enforces_ownership_and_preserves_clickhouse_envelope( + recording_server: RecordingServer, role: str, user_id: str | None, expected_status: int ) -> None: from fastapi import FastAPI from fastapi.testclient import TestClient from litellm.proxy._types import UserAPIKeyAuth + from litellm.proxy.auth.authorization_dependencies import get_log_team_lookup from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.tracing_endpoints import provide_receiver, provide_trace_query_secret, router @@ -312,19 +322,28 @@ def test_trace_sql_endpoint_executes_for_admin_and_preserves_clickhouse_envelope "rows": 1, "statistics": {"elapsed": 0.01, "rows_read": 1, "bytes_read": 1}, } - recording_server.expected_requests = 12 - for _ in range(11): - recording_server.enqueue(ResponseSpec(body="")) - recording_server.enqueue(ResponseSpec(body=envelope)) + recording_server.expected_requests = 12 if expected_status == 200 else 0 + if expected_status == 200: + for _ in range(11): + recording_server.enqueue(ResponseSpec(body="")) + recording_server.enqueue(ResponseSpec(body=envelope)) storage: Final = ClickHouseStorage(TraceStorageConfig(recording_server.base_url, "trace_test")) app: Final = FastAPI() app.include_router(router) app.dependency_overrides[provide_trace_query_secret] = lambda: "test-master-secret" - app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=role, token="test") + app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=role, user_id=user_id, token="test") app.dependency_overrides[provide_receiver] = lambda: TraceReceiver(TraceStore(storage)) + + async def permitted_teams(auth: UserAPIKeyAuth) -> tuple[str, ...]: + return () + + app.dependency_overrides[get_log_team_lookup] = lambda: permitted_teams with TestClient(app) as client: result: Final = client.post("/v1/traces/query", json={"sql": "SELECT 42 AS answer"}) - assert result.status_code == 200, result.text + assert result.status_code == expected_status, result.text + if expected_status == 403: + assert result.json() == {"detail": "Not allowed to view logs"} + return assert result.json() == envelope assert recording_server.requests[-1].raw_body == b"SELECT 42 AS answer" assert client.post("/v1/traces/query", json={"sql": " "}).status_code == 400 diff --git a/tests/unit/proxy/auth/test_authorization.py b/tests/unit/proxy/auth/test_authorization.py new file mode 100644 index 00000000000..7d1548dd828 --- /dev/null +++ b/tests/unit/proxy/auth/test_authorization.py @@ -0,0 +1,29 @@ +from typing import Final + +import pytest + +from litellm.proxy._types import UserAPIKeyAuth +from litellm.proxy.auth.authorization import OwnedRows, resolve_owned_read_scope, resolve_trace_read_scope + + +@pytest.mark.asyncio +@pytest.mark.parametrize("token", (None, "key")) +async def test_team_membership_or_key_without_user_does_not_grant_log_access(token: str | None) -> None: + async def unexpected_lookup() -> tuple[str, ...]: + pytest.fail("Identity-less callers cannot consult team permissions") + + assert await resolve_trace_read_scope(UserAPIKeyAuth(team_id="team", token=token), unexpected_lookup) is None + + +@pytest.mark.asyncio +@pytest.mark.parametrize("token", (None, "key")) +@pytest.mark.parametrize("lookup_fails", (False, True)) +async def test_trace_reads_share_user_and_team_scope_regardless_of_key(token: str | None, lookup_fails: bool) -> None: + async def lookup() -> tuple[str, ...]: + if lookup_fails: + raise RuntimeError("team lookup failed") + return ("permitted",) + + expected: Final = OwnedRows("caller", () if lookup_fails else ("permitted",)) + assert await resolve_owned_read_scope("caller", lookup) == expected + assert await resolve_trace_read_scope(UserAPIKeyAuth(user_id="caller", token=token), lookup) == expected diff --git a/tests/unit/proxy/common_utils/test_http_parsing_utils.py b/tests/unit/proxy/common_utils/test_http_parsing_utils.py index fd747d5a6f2..f00be0f8a65 100644 --- a/tests/unit/proxy/common_utils/test_http_parsing_utils.py +++ b/tests/unit/proxy/common_utils/test_http_parsing_utils.py @@ -1354,7 +1354,7 @@ async def test_auth_body_read_and_trace_handler_leave_stream_for_receiver_limit( store: Final = MagicMock() store.insert_spans = AsyncMock() context: Final = await tracing_endpoints.provide_trace_access( - auth=UserAPIKeyAuth(token="key", team_id="team"), tracing=TraceReceiver(store) + auth=UserAPIKeyAuth(token="key", team_id="team"), tracing=TraceReceiver(store), log_team_lookup=AsyncMock() ) parsed, parse_error = await _read_request_body_deferring_parse_failure(request) diff --git a/tests/unit/proxy/management_endpoints/management_v1/test_spend_logs.py b/tests/unit/proxy/management_endpoints/management_v1/test_spend_logs.py index b6867d338c5..7523c864985 100644 --- a/tests/unit/proxy/management_endpoints/management_v1/test_spend_logs.py +++ b/tests/unit/proxy/management_endpoints/management_v1/test_spend_logs.py @@ -1,5 +1,5 @@ from datetime import datetime, timezone -from unittest.mock import AsyncMock, MagicMock, patch +from unittest.mock import AsyncMock, MagicMock import pytest from fastapi import FastAPI, Request @@ -7,6 +7,7 @@ from fastapi.exceptions import RequestValidationError from fastapi.testclient import TestClient from litellm.proxy._types import LiteLLMRoutes, LitellmUserRoles +from litellm.proxy.auth.authorization_dependencies import get_log_team_lookup from litellm.proxy.auth.user_api_key_auth import UserAPIKeyAuth, user_api_key_auth from litellm.proxy.list_api.common import ( PROBLEM_TYPE_BASE, @@ -51,7 +52,7 @@ WINDOW = "filter[startTime][gte]=2026-07-23T00:00:00Z&filter[startTime][lte]=202 @pytest.fixture def mock_prisma_client(monkeypatch): prisma_client = MagicMock() - prisma_client.db.query_raw = AsyncMock(return_value=[]) + prisma_client.db.query_raw = AsyncMock(return_value=()) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma_client) return prisma_client @@ -71,9 +72,10 @@ def _mock_rows(mock_prisma_client, end_users: list[str]) -> AsyncMock: return query_raw -def _as_role(role: LitellmUserRoles, user_id): +def _as_role(role: LitellmUserRoles, user_id, log_team_lookup): original = app.dependency_overrides.copy() app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth(user_id=user_id, user_role=role) + app.dependency_overrides[get_log_team_lookup] = lambda: log_team_lookup return original @@ -283,13 +285,9 @@ def test_applies_no_scope_for_a_proxy_admin(mock_prisma_client, as_proxy_admin): def test_scopes_a_team_admin_to_their_own_rows_and_teams(mock_prisma_client, role): """A team admin must not see end users belonging to teams they cannot read.""" query_raw = _mock_rows(mock_prisma_client, ["cust-a"]) - original = _as_role(role, user_id="team-admin-1") + original = _as_role(role, user_id="team-admin-1", log_team_lookup=AsyncMock(return_value=("team-a", "team-b"))) try: - with patch( - "litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs", - new=AsyncMock(return_value=["team-a", "team-b"]), - ): - response = _get() + response = _get() finally: app.dependency_overrides = original @@ -297,24 +295,20 @@ def test_scopes_a_team_admin_to_their_own_rows_and_teams(mock_prisma_client, rol # Same clause shape ui_view_spend_logs builds, so the two cannot diverge. assert '("user" = $3 OR team_id = ANY($4::text[]))' in query_raw.call_args.args[0] assert query_raw.call_args.args[3] == "team-admin-1" - assert query_raw.call_args.args[4] == ["team-a", "team-b"] + assert query_raw.call_args.args[4] == ("team-a", "team-b") def test_scopes_a_teamless_user_to_their_own_rows(mock_prisma_client): query_raw = _mock_rows(mock_prisma_client, []) - original = _as_role(LitellmUserRoles.INTERNAL_USER, user_id="solo") + original = _as_role(LitellmUserRoles.INTERNAL_USER, user_id="solo", log_team_lookup=AsyncMock(return_value=())) try: - with patch( - "litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs", - new=AsyncMock(return_value=[]), - ): - response = _get() + response = _get() finally: app.dependency_overrides = original assert response.status_code == 200 sql = query_raw.call_args.args[0] - assert '("user" = $3)' in sql + assert '"user" = $3' in sql assert "team_id" not in sql assert query_raw.call_args.args[3] == "solo" @@ -322,13 +316,9 @@ def test_scopes_a_teamless_user_to_their_own_rows(mock_prisma_client): def test_returns_nothing_when_the_caller_owns_no_scope(mock_prisma_client): """Unidentifiable caller must match no rows, never fall through to unscoped.""" query_raw = _mock_rows(mock_prisma_client, []) - original = _as_role(LitellmUserRoles.INTERNAL_USER, user_id=None) + original = _as_role(LitellmUserRoles.INTERNAL_USER, user_id=None, log_team_lookup=AsyncMock(return_value=())) try: - with patch( - "litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs", - new=AsyncMock(return_value=[]), - ): - response = _get() + response = _get() finally: app.dependency_overrides = original @@ -339,19 +329,17 @@ def test_returns_nothing_when_the_caller_owns_no_scope(mock_prisma_client): def test_scopes_when_the_permitted_team_lookup_fails(mock_prisma_client): """A failed team lookup must degrade to own-rows-only, never to unscoped.""" query_raw = _mock_rows(mock_prisma_client, []) - original = _as_role(LitellmUserRoles.INTERNAL_USER, user_id="solo") + original = _as_role( + LitellmUserRoles.INTERNAL_USER, user_id="solo", log_team_lookup=AsyncMock(side_effect=RuntimeError("db down")) + ) try: - with patch( - "litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs", - new=AsyncMock(side_effect=RuntimeError("db down")), - ): - response = _get() + response = _get() finally: app.dependency_overrides = original assert response.status_code == 200 sql = query_raw.call_args.args[0] - assert '("user" = $3)' in sql + assert '"user" = $3' in sql assert "team_id" not in sql @@ -422,24 +410,22 @@ def test_user_facet_reads_internal_users_from_spend_logs(mock_prisma_client, as_ def test_user_facet_uses_the_same_team_scope_as_request_logs(mock_prisma_client): query_raw = AsyncMock(return_value=[{"user": "member@example.com"}]) mock_prisma_client.db.query_raw = query_raw - original = _as_role(LitellmUserRoles.INTERNAL_USER, user_id="team-admin-1") + original = _as_role( + LitellmUserRoles.INTERNAL_USER, user_id="team-admin-1", log_team_lookup=AsyncMock(return_value=("team-a",)) + ) try: - with patch( - "litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs", - new=AsyncMock(return_value=["team-a"]), - ): - response = _get_users() + response = _get_users() finally: app.dependency_overrides = original assert response.status_code == 200 assert '("user" = $3 OR team_id = ANY($4::text[]))' in query_raw.call_args.args[0] assert query_raw.call_args.args[3] == "team-admin-1" - assert query_raw.call_args.args[4] == ["team-a"] + assert query_raw.call_args.args[4] == ("team-a",) def test_user_facet_searches_the_internal_user_value(mock_prisma_client, as_proxy_admin): - query_raw = AsyncMock(return_value=[]) + query_raw = AsyncMock(return_value=()) mock_prisma_client.db.query_raw = query_raw _get_users(f"{WINDOW}&q=alice%40example.com") diff --git a/tests/unit/proxy/spend_tracking/test_log_visibility.py b/tests/unit/proxy/spend_tracking/test_log_visibility.py deleted file mode 100644 index 140225d2000..00000000000 --- a/tests/unit/proxy/spend_tracking/test_log_visibility.py +++ /dev/null @@ -1,47 +0,0 @@ -from typing import Final - -import pytest -from fastapi import HTTPException - -from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth -from litellm.proxy.spend_tracking.log_visibility import LogVisibility, log_visibility - - -@pytest.mark.asyncio -@pytest.mark.parametrize( - ("auth", "expected"), - ( - (UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), LogVisibility(all_teams=True)), - (UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY), LogVisibility(all_teams=True)), - ( - UserAPIKeyAuth(user_id="user", token="key", team_id="unpermitted"), - LogVisibility(user_id="user", team_ids=("permitted",), api_key_hash="key"), - ), - (UserAPIKeyAuth(token="key", team_id="unpermitted"), LogVisibility(api_key_hash="key")), - ), -) -async def test_log_visibility_uses_user_and_permitted_teams_instead_of_key_team_membership( - auth: UserAPIKeyAuth, - expected: LogVisibility, -) -> None: - async def permitted_teams(caller: UserAPIKeyAuth) -> tuple[str, ...]: - assert caller is auth - return ("permitted",) - - assert await log_visibility(auth, permitted_teams) == expected - - -@pytest.mark.asyncio -async def test_missing_team_permissions_preserve_authenticated_user_and_key_visibility() -> None: - async def no_teams(auth: UserAPIKeyAuth) -> tuple[str, ...]: - return () - - auth: Final = UserAPIKeyAuth(user_id="user", token="key", team_id="team") - assert await log_visibility(auth, no_teams) == LogVisibility(user_id=auth.user_id, api_key_hash="key") - - -@pytest.mark.asyncio -async def test_team_membership_without_authenticated_identity_does_not_grant_log_access() -> None: - with pytest.raises(HTTPException) as error: - await log_visibility(UserAPIKeyAuth(team_id="team")) - assert error.value.status_code == 403 diff --git a/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py b/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py index 506de58e438..6aaa536a945 100644 --- a/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py +++ b/tests/unit/proxy/spend_tracking/test_spend_management_endpoints.py @@ -14,6 +14,8 @@ from fastapi.testclient import TestClient import litellm import litellm.proxy.proxy_server as ps +from litellm.proxy.auth.authorization import OwnedRows +from litellm.proxy.auth.authorization_dependencies import get_log_team_lookup, load_permitted_log_team_ids def _default_date_range(): @@ -1656,10 +1658,7 @@ async def test_ui_view_spend_logs_explicit_user_filter_cannot_escape_own_scope(c "litellm.proxy.proxy_server.prisma_client", make_ui_spend_logs_mock_prisma([caller_log], lambda _where: [], query_observer=observe_query), ) - monkeypatch.setattr( - "litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs", - AsyncMock(return_value=[]), - ) + monkeypatch.setitem(app.dependency_overrides, get_log_team_lookup, lambda: AsyncMock(return_value=())) app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( user_role=LitellmUserRoles.INTERNAL_USER, user_id="caller@example.com" ) @@ -1715,10 +1714,7 @@ async def test_ui_view_spend_logs_without_user_filter_includes_permitted_team_sc "litellm.proxy.proxy_server.prisma_client", make_ui_spend_logs_mock_prisma([caller_log, member_log, outside_log], filter_by_scope), ) - monkeypatch.setattr( - "litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs", - AsyncMock(return_value=["team-9"]), - ) + monkeypatch.setitem(app.dependency_overrides, get_log_team_lookup, lambda: AsyncMock(return_value=("team-9",))) app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( user_role=LitellmUserRoles.INTERNAL_USER, user_id="team-admin@example.com" ) @@ -1738,21 +1734,13 @@ async def test_ui_view_spend_logs_without_user_filter_includes_permitted_team_sc @pytest.mark.asyncio -async def test_permitted_team_scope_falls_back_to_own_user_when_lookup_fails(monkeypatch): - monkeypatch.setattr( - "litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs", - AsyncMock(side_effect=RuntimeError("database unavailable")), - ) +async def test_permitted_team_scope_falls_back_to_own_user_when_lookup_fails(): + from litellm.proxy.auth.authorization import resolve_owned_read_scope - permitted_team_ids = await spend_management_endpoints._get_permitted_team_ids_for_spend_logs_or_empty( - prisma_client=MagicMock(), - user_api_key_dict=UserAPIKeyAuth( - user_role=LitellmUserRoles.INTERNAL_USER, - user_id="caller@example.com", - ), - ) + async def unavailable(): + raise RuntimeError("database unavailable") - assert permitted_team_ids == () + assert await resolve_owned_read_scope("caller", unavailable) == OwnedRows("caller") @pytest.mark.asyncio @@ -1876,10 +1864,7 @@ async def test_ui_view_spend_logs_user_filter_intersects_permitted_team_scope(cl "litellm.proxy.proxy_server.prisma_client", make_ui_spend_logs_mock_prisma([member_log, other_team_log], filter_by_user_and_scope), ) - monkeypatch.setattr( - "litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs", - AsyncMock(return_value=["team-9"]), - ) + monkeypatch.setitem(app.dependency_overrides, get_log_team_lookup, lambda: AsyncMock(return_value=("team-9",))) app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( user_role=LitellmUserRoles.INTERNAL_USER, user_id="team-admin" ) @@ -2129,61 +2114,6 @@ async def test_ui_view_session_spend_logs_rehydrates_metadata_jsonb_text(client, app.dependency_overrides.pop(ps.user_api_key_auth, None) -@pytest.mark.asyncio -async def test_ui_view_session_spend_logs_scopes_non_admin_to_own_logs(client, monkeypatch): - own_log = { - "id": "log1", - "request_id": "req1", - "session_id": "session-123", - "user": "user-1", - "startTime": "2024-01-01T00:00:00Z", - } - - class MockDB: - async def count(self, *args, **kwargs): - assert kwargs.get("where") == {"session_id": "session-123", "user": "user-1"} - return 1 - - async def query_raw(self, sql_query, session_id, page_size, skip, scoped_user): - assert session_id == "session-123" - assert scoped_user == "user-1" - assert '"user" = $4' in sql_query - return [own_log] - - class MockPrismaClient: - def __init__(self): - self.db = MockDB() - self.db.litellm_spendlogs = self.db - - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", MockPrismaClient()) - - async def no_permitted_teams(*args, **kwargs): - return [] - - monkeypatch.setattr( - "litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs", - no_permitted_teams, - ) - - app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( - user_role=LitellmUserRoles.INTERNAL_USER, user_id="user-1" - ) - - try: - response = client.get( - "/spend/logs/session/ui", - params={"session_id": "session-123", "page": 1, "page_size": 50}, - headers={"Authorization": "Bearer sk-test"}, - ) - - assert response.status_code == 200 - data = response.json() - assert data["total"] == 1 - assert [row["request_id"] for row in data["data"]] == ["req1"] - finally: - app.dependency_overrides.pop(ps.user_api_key_auth, None) - - @pytest.mark.asyncio async def test_ui_view_session_spend_logs_includes_permitted_team_logs(client, monkeypatch): class MockDB: @@ -2200,7 +2130,7 @@ async def test_ui_view_session_spend_logs_includes_permitted_team_logs(client, m async def query_raw(self, sql_query, session_id, page_size, skip, scoped_user, team_ids): assert session_id == "session-123" assert scoped_user == "user-1" - assert team_ids == ["team-9"] + assert tuple(team_ids) == ("team-9",) assert '("user" = $4 OR team_id = ANY($5::text[]))' in sql_query return [ { @@ -2222,10 +2152,7 @@ async def test_ui_view_session_spend_logs_includes_permitted_team_logs(client, m async def permitted_teams(*args, **kwargs): return ["team-9"] - monkeypatch.setattr( - "litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs", - permitted_teams, - ) + monkeypatch.setitem(app.dependency_overrides, get_log_team_lookup, lambda: permitted_teams) app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( user_role=LitellmUserRoles.INTERNAL_USER, user_id="user-1" @@ -2658,31 +2585,6 @@ async def test_ui_view_spend_logs_request_id_rejects_foreign_row_inserted_after_ app.dependency_overrides.pop(ps.user_api_key_auth, None) -def _make_payload_lookup_prisma(rows): - """Emulate the detail endpoint's SQL over an in-memory corpus: the owner - pre-check, the caller scope on ``"user"`` and permitted teams, and the - exact-request_id-first ordering with LIMIT 1.""" - - class MockDB: - async def query_raw(self, sql_query, *params): - if 'SELECT DISTINCT "user", team_id' in sql_query: - return _emulate_spend_log_owner_lookup(rows, sql_query, params) - lookup_id = params[0] - matches = [r for r in rows if lookup_id in (r["request_id"], r["litellm_call_id"])] - if '"user" = $2' in sql_query: - team_ids = params[2] if "ANY($3::text[])" in sql_query else () - matches = [r for r in matches if r["user"] == params[1] or r["team_id"] in team_ids] - if "ORDER BY (request_id = $1) DESC" in sql_query: - matches = sorted(matches, key=lambda r: r["request_id"] == lookup_id, reverse=True) - return matches[:1] - - class MockPrisma: - def __init__(self): - self.db = MockDB() - - return MockPrisma() - - def _payload_row(request_id, litellm_call_id, user, prompt): return { "request_id": request_id, @@ -2696,36 +2598,6 @@ def _payload_row(request_id, litellm_call_id, user, prompt): } -@pytest.mark.asyncio -async def test_ui_view_request_response_collision_serves_callers_own_row(client, monkeypatch): - """The attacker's row carries the victim's request_id as its client-set call id - and was written first. Each tenant's detail lookup of that id serves only their - own payload, and an admin's lookup resolves the exact request_id match rather - than whichever colliding row the database happens to return first.""" - prisma = _make_payload_lookup_prisma( - [ - _payload_row("attacker-req", "victim-req", "attacker_user", "attacker prompt"), - _payload_row("victim-req", "victim-call-id", "victim_user", "victim prompt"), - ] - ) - monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", prisma) - try: - for role, user_id, own_prompt, other_prompt in ( - (LitellmUserRoles.INTERNAL_USER, "victim_user", "victim prompt", "attacker prompt"), - (LitellmUserRoles.INTERNAL_USER, "attacker_user", "attacker prompt", "victim prompt"), - (LitellmUserRoles.PROXY_ADMIN, "admin", "victim prompt", "attacker prompt"), - ): - app.dependency_overrides[ps.user_api_key_auth] = lambda role=role, user_id=user_id: UserAPIKeyAuth( - user_role=role, user_id=user_id - ) - response = client.get("/spend/logs/ui/victim-req", headers={"Authorization": "Bearer sk-test"}) - assert response.status_code == 200, response.text - assert own_prompt in response.text - assert other_prompt not in response.text - finally: - app.dependency_overrides.pop(ps.user_api_key_auth, None) - - @pytest.mark.asyncio async def test_ui_view_request_response_rejects_foreign_row_inserted_after_owner_check(client, monkeypatch): """Backstop behind the SQL scope on the detail endpoint (the mock ignores the @@ -2818,11 +2690,15 @@ async def test_ui_view_request_response_custom_logger_is_keyed_by_callers_own_re that id as its request_id. The custom logger is asked for the caller's own stored request_id, so the caller gets their payload rather than a 403 from the foreign payload's owner check, and the foreign payload is never fetched.""" - prisma = _make_payload_lookup_prisma( - [ - _payload_row("shared-id", "other-call-id", "other_user", "other tenant prompt"), - _payload_row("caller-req", "shared-id", "caller_user", "caller prompt"), - ] + prisma = MagicMock( + db=MagicMock( + query_raw=AsyncMock( + side_effect=[ + [{"user": "other_user", "team_id": None}, {"user": "caller_user", "team_id": None}], + [_payload_row("caller-req", "shared-id", "caller_user", "caller prompt")], + ] + ) + ) ) cold_storage = { "shared-id": { @@ -3161,10 +3037,7 @@ async def test_ui_view_spend_logs_search_keeps_non_admin_scope(client, monkeypat "litellm.proxy.proxy_server.prisma_client", make_ui_spend_logs_mock_prisma(logs, _search_filter_fn(logs, captured)), ) - monkeypatch.setattr( - "litellm.proxy.spend_tracking.spend_management_endpoints._get_permitted_team_ids_for_spend_logs", - AsyncMock(return_value=[]), - ) + monkeypatch.setitem(app.dependency_overrides, get_log_team_lookup, lambda: AsyncMock(return_value=())) ownership_check = AsyncMock() monkeypatch.setattr( "litellm.proxy.spend_tracking.spend_management_endpoints._assert_user_can_view_request_id", @@ -3405,9 +3278,7 @@ async def test_ui_view_spend_logs_with_used_client_oauth_token_filter(client, mo start_date, end_date = _default_date_range() - app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth( - user_role=LitellmUserRoles.PROXY_ADMIN - ) + app.dependency_overrides[ps.user_api_key_auth] = lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN) try: for flag, expected_ids in (("true", ["req-seat"]), ("false", ["req-key"])): response = client.get( @@ -7906,3 +7777,115 @@ def test_capture_rate_reports_an_unreadable_bill_as_502(client, monkeypatch): app.dependency_overrides.pop(ps.user_api_key_auth, None) assert response.status_code == 502 assert "HTTP 401" in response.json()["detail"] + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("user_id", "owner_user", "owner_team", "permitted", "expected"), + [ + ("caller", "caller", "broken", False, True), + ("caller", "other", "allowed", True, True), + ("caller", "other", "allowed", False, False), + ("caller", "other", None, True, False), + (None, None, None, True, False), + (None, None, "allowed", True, True), + ], +) +async def test_shared_owner_policy_preserves_own_user_and_team_access( + user_id, owner_user, owner_team, permitted, expected +): + from litellm.proxy.auth.authorization import can_read_log_owner + + async def lookup(team_id): + if team_id == "broken": + raise RuntimeError("team lookup failed") + return permitted + + assert await can_read_log_owner(user_id, owner_user, owner_team, lookup) is expected + + +@pytest.mark.asyncio +async def test_shared_owner_policy_propagates_team_lookup_failure(): + from litellm.proxy.auth.authorization import can_read_log_owner + + async def unavailable(team_id): + raise RuntimeError("team lookup failed") + + with pytest.raises(RuntimeError, match="team lookup failed"): + await can_read_log_owner("caller", "other", "team", unavailable) + + +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("params", "expected_status"), + [ + ({"start_date": "invalid", "end_date": "invalid"}, 400), + ({"request_id": "foreign"}, 403), + ], +) +async def test_log_team_dependency_preserves_checks_before_permission_lookup( + client, monkeypatch, params, expected_status +): + from litellm.proxy._types import LiteLLM_UserTable + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + team_reads = [] + + class TeamTable: + async def find_many(self, where): + team_reads.append(where) + return [] + + cache = UserApiKeyCache() + await cache.async_set_cache( + key="caller", value=LiteLLM_UserTable(user_id="caller", teams=["team"]), model_type=LiteLLM_UserTable + ) + prisma = MagicMock( + db=MagicMock( + query_raw=AsyncMock(return_value=[{"user": "other", "team_id": None}]), + litellm_teamtable=TeamTable(), + ) + ) + monkeypatch.setattr(ps, "prisma_client", prisma) + monkeypatch.setattr(ps, "user_api_key_cache", cache) + monkeypatch.setitem( + app.dependency_overrides, + ps.user_api_key_auth, + lambda: UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER, user_id="caller"), + ) + + response = client.get("/spend/logs/ui", params=params, headers={"Authorization": "Bearer sk-test"}) + + assert response.status_code == expected_status, response.text + assert team_reads == [] + + +@pytest.mark.asyncio +async def test_management_team_lookup_without_memberships_keeps_own_user_scope(): + from litellm.proxy._types import LiteLLM_UserTable + from litellm.proxy.auth.authorization import resolve_owned_read_scope + from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache + + cache = UserApiKeyCache() + await cache.async_set_cache( + key="caller", value=LiteLLM_UserTable(user_id="caller", teams=[]), model_type=LiteLLM_UserTable + ) + auth = UserAPIKeyAuth(user_id="caller", user_role=LitellmUserRoles.INTERNAL_USER) + team_reads = [] + + class TeamTable: + async def find_many(self, where): + team_reads.append(where) + return [] + + prisma = MagicMock(db=MagicMock(litellm_teamtable=TeamTable())) + + async def lookup(): + return await load_permitted_log_team_ids( + auth, prisma_client=prisma, user_api_key_cache=cache, proxy_logging_obj=ps.proxy_logging_obj + ) + + assert await lookup() == () + scope = await resolve_owned_read_scope(auth.user_id, lookup) + assert scope == OwnedRows("caller") + assert team_reads == [] diff --git a/tests/unit/proxy/spend_tracking/test_spend_query_optimization.py b/tests/unit/proxy/spend_tracking/test_spend_query_optimization.py index 6752c91e9f2..93fae093340 100644 --- a/tests/unit/proxy/spend_tracking/test_spend_query_optimization.py +++ b/tests/unit/proxy/spend_tracking/test_spend_query_optimization.py @@ -11,7 +11,7 @@ from unittest.mock import AsyncMock, MagicMock import pytest - +from litellm.proxy.auth.authorization_dependencies import get_log_team_lookup from litellm.proxy.spend_tracking.spend_tracking_utils import ( get_spend_by_team, get_spend_by_team_and_customer, @@ -180,6 +180,7 @@ async def test_spend_logs_ui_wraps_params_in_at_time_zone_utc(monkeypatch): mock_request.url.path = "/spend/logs/ui" await ui_view_spend_logs( + log_team_lookup=await get_log_team_lookup(), request=mock_request, api_key=None, user_id=None, @@ -209,9 +210,7 @@ def _make_ui_spend_logs_mock(count_total, page_rows): """ mock_prisma = MagicMock() mock_prisma.db = MagicMock() - mock_prisma.db.query_raw = AsyncMock( - side_effect=[[{"total_count": count_total}], page_rows] - ) + mock_prisma.db.query_raw = AsyncMock(side_effect=[[{"total_count": count_total}], page_rows]) mock_prisma.db.litellm_spendlogs = MagicMock() mock_prisma.db.litellm_spendlogs.count = AsyncMock(return_value=0) return mock_prisma @@ -244,6 +243,7 @@ async def test_spend_logs_ui_uses_bounded_count_not_full_scan(monkeypatch): mock_request.url.path = "/spend/logs/ui" response = await ui_view_spend_logs( + log_team_lookup=await get_log_team_lookup(), request=mock_request, api_key=None, user_id=None, @@ -264,17 +264,13 @@ async def test_spend_logs_ui_uses_bounded_count_not_full_scan(monkeypatch): count_sql = count_call[0][0] assert "COUNT(*) OVER ()" not in count_sql assert "LIMIT" in count_sql and "FROM (" in count_sql, ( - "the total must come from a bounded subquery count, not a full-window " - f"scan. SQL was:\n{count_sql}" - ) - assert count_call[0][-1] == SPEND_LOGS_PAGINATION_COUNT_CAP + 1, ( - "the bounded count must probe at most cap+1 rows" + f"the total must come from a bounded subquery count, not a full-window scan. SQL was:\n{count_sql}" ) + assert count_call[0][-1] == SPEND_LOGS_PAGINATION_COUNT_CAP + 1, "the bounded count must probe at most cap+1 rows" page_sql = mock_prisma.db.query_raw.call_args_list[1][0][0] assert "COUNT(*) OVER ()" not in page_sql, ( - "the page query must not carry a window count that forces a full-window " - f"scan. SQL was:\n{page_sql}" + f"the page query must not carry a window count that forces a full-window scan. SQL was:\n{page_sql}" ) assert "GROUP BY" not in count_sql and "DISTINCT ON" not in page_sql, ( "without group_by_session the endpoint must keep raw per-call pagination" @@ -302,9 +298,7 @@ async def test_spend_logs_ui_caps_total_for_large_result_sets(monkeypatch): ) page_rows = [{"request_id": "req-1", "metadata": "{}", "session_id": None}] - mock_prisma = _make_ui_spend_logs_mock( - count_total=SPEND_LOGS_PAGINATION_COUNT_CAP + 1, page_rows=page_rows - ) + mock_prisma = _make_ui_spend_logs_mock(count_total=SPEND_LOGS_PAGINATION_COUNT_CAP + 1, page_rows=page_rows) monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", mock_prisma) auth = UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN, user_id="admin") @@ -312,6 +306,7 @@ async def test_spend_logs_ui_caps_total_for_large_result_sets(monkeypatch): mock_request.url.path = "/spend/logs/ui" response = await ui_view_spend_logs( + log_team_lookup=await get_log_team_lookup(), request=mock_request, api_key=None, user_id=None, @@ -358,6 +353,7 @@ async def test_spend_logs_ui_empty_page_reports_zero_total(monkeypatch): mock_request.url.path = "/spend/logs/ui" response = await ui_view_spend_logs( + log_team_lookup=await get_log_team_lookup(), request=mock_request, api_key=None, user_id=None, @@ -406,6 +402,7 @@ async def test_spend_logs_ui_out_of_range_page_keeps_total(monkeypatch): mock_request.url.path = "/spend/logs/ui" response = await ui_view_spend_logs( + log_team_lookup=await get_log_team_lookup(), request=mock_request, api_key=None, user_id=None, @@ -553,6 +550,7 @@ async def test_spend_logs_ui_group_by_session_paginates_sessions(monkeypatch): mock_request.url.path = "/spend/logs/ui" response = await ui_view_spend_logs( + log_team_lookup=await get_log_team_lookup(), request=mock_request, api_key=None, user_id=None, @@ -620,6 +618,7 @@ async def test_spend_logs_ui_group_by_session_offset_pages_for_other_sorts(monke mock_request.url.path = "/spend/logs/ui" response = await ui_view_spend_logs( + log_team_lookup=await get_log_team_lookup(), request=mock_request, api_key=None, user_id=None, @@ -669,6 +668,7 @@ async def test_spend_logs_ui_request_id_lookup_with_grouping_returns_exact_row(m mock_request.url.path = "/spend/logs/ui" response = await ui_view_spend_logs( + log_team_lookup=await get_log_team_lookup(), request=mock_request, api_key=None, user_id=None, diff --git a/tests/unit/proxy/test_tracing_endpoints.py b/tests/unit/proxy/test_tracing_endpoints.py index fd05d54dcb6..d901222aaf4 100644 --- a/tests/unit/proxy/test_tracing_endpoints.py +++ b/tests/unit/proxy/test_tracing_endpoints.py @@ -5,7 +5,7 @@ Tests for the agent tracing endpoints (litellm/proxy/tracing_endpoints.py). from collections.abc import AsyncGenerator, Mapping from contextlib import asynccontextmanager from types import ModuleType -from typing import Final +from typing import Final, Literal from unittest.mock import AsyncMock, MagicMock import pytest @@ -14,12 +14,14 @@ from fastapi.testclient import TestClient from litellm.proxy import tracing_endpoints from litellm.proxy._types import LitellmUserRoles, ProxyLifespanState, UserAPIKeyAuth +from litellm.proxy.auth.authorization import OwnedRows, ReadScope +from litellm.proxy.auth.authorization_dependencies import get_log_team_lookup from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.tracing_runtime import manage_tracing, provide_storage from litellm.rust_bridge import loader from litellm.rust_bridge.trace_queries import SPAN_DETAIL, SpanDetailParams from litellm.rust_bridge.trace_query_responses import TraceQueryHelp, TraceSQLResponse -from litellm.rust_bridge.traces import AdminQueryScope, ClickHouseStorage, TraceStorageConfig +from litellm.rust_bridge.traces import AllQueryScope, ClickHouseStorage, TraceStorageConfig from litellm.tracing import TraceReceiver, TracingPayloadTooLargeError from litellm.tracing.store import TraceStore from litellm.tracing.types import TraceScope @@ -57,7 +59,11 @@ QUERY_HELP: Final[Mapping[str, object]] = { TEAM_KEY = UserAPIKeyAuth( - token="hashed-key", team_id="team-research", org_id="org-1", user_role=LitellmUserRoles.INTERNAL_USER + user_id="user", + token="hashed-key", + team_id="team-research", + org_id="org-1", + user_role=LitellmUserRoles.INTERNAL_USER, ) TRACE_RESPONSE: Final = { "summary": { @@ -97,38 +103,47 @@ SPAN_DETAIL_RESPONSE: Final = { ( pytest.param( UserAPIKeyAuth(token="admin-key", team_id="team-a", user_role=LitellmUserRoles.PROXY_ADMIN), - TraceScope(all_teams=1, user_id="", team_ids=(), api_key_hash=""), + TraceScope(all_teams=1, user_id="", team_ids=()), True, id="admin", ), pytest.param( UserAPIKeyAuth(token="view-key", team_id="team-a", user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY), - TraceScope(all_teams=1, user_id="", team_ids=(), api_key_hash=""), + TraceScope(all_teams=1, user_id="", team_ids=()), False, id="view-only-admin", ), pytest.param( TEAM_KEY, - TraceScope(all_teams=0, user_id="", team_ids=(), api_key_hash="hashed-key"), + TraceScope(all_teams=0, user_id="user", team_ids=()), True, id="team-key", ), pytest.param( - UserAPIKeyAuth(token="hashed-key", user_role=LitellmUserRoles.INTERNAL_USER), - TraceScope(all_teams=0, user_id="", team_ids=(), api_key_hash="hashed-key"), + UserAPIKeyAuth(user_id="user", token="hashed-key", user_role=LitellmUserRoles.INTERNAL_USER), + TraceScope(all_teams=0, user_id="user", team_ids=()), True, id="teamless-key", ), + pytest.param( + UserAPIKeyAuth(token="hashed-key", user_role=LitellmUserRoles.INTERNAL_USER), + None, + True, + id="key-without-user-can-only-write", + ), ), ) def test_trace_read_and_write_permissions( - client: TestClient, receiver: MagicMock, auth: UserAPIKeyAuth, scope: TraceScope, can_write: bool + client: TestClient, receiver: MagicMock, auth: UserAPIKeyAuth, scope: TraceScope | None, can_write: bool ) -> None: client.app.dependency_overrides[user_api_key_auth] = lambda: auth read: Final = client.get("/v1/traces?start_ms=1&end_ms=2") - assert read.status_code == 200, read.text - receiver.list_traces.assert_awaited_once_with(scope=scope, start_ms=1, end_ms=2, cursor=None) + assert read.status_code == (403 if scope is None else 200), read.text + if scope is None: + receiver.list_traces.assert_not_awaited() + else: + receiver.list_traces.assert_awaited_once_with(scope=scope, start_ms=1, end_ms=2, cursor=None) write: Final = client.post("/v1/traces", json={}) assert write.status_code == (200 if can_write else 403), write.text @@ -160,6 +175,11 @@ def client() -> TestClient: app = FastAPI() app.include_router(tracing_endpoints.router) app.dependency_overrides[user_api_key_auth] = lambda: TEAM_KEY + + async def lookup(auth: UserAPIKeyAuth) -> tuple[str, ...]: + return () + + app.dependency_overrides[get_log_team_lookup] = lambda: lookup return TestClient(app) @@ -227,7 +247,7 @@ def test_list_traces_passes_scope_window_and_cursor(client, receiver): assert response.status_code == 200 assert response.json() == {"data": [], "next_cursor": None} receiver.list_traces.assert_awaited_once_with( - scope={"all_teams": 0, "user_id": "", "team_ids": (), "api_key_hash": "hashed-key"}, + scope={"all_teams": 0, "user_id": "user", "team_ids": ()}, start_ms=1, end_ms=2, cursor="abc", @@ -247,9 +267,7 @@ def test_get_trace_404_and_200(client, receiver): response = client.get("/v1/traces/t1") assert response.status_code == 200 assert response.json() == TRACE_RESPONSE - receiver.get_trace.assert_awaited_with( - "t1", {"all_teams": 0, "user_id": "", "team_ids": (), "api_key_hash": "hashed-key"}, "" - ) + receiver.get_trace.assert_awaited_with("t1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "") def test_get_span_404_and_200(client, receiver): @@ -258,9 +276,7 @@ def test_get_span_404_and_200(client, receiver): response = client.get("/v1/traces/t1/spans/s1") assert response.status_code == 200 assert response.json()["span_id"] == "s1" - receiver.get_span.assert_awaited_with( - "t1", "s1", {"all_teams": 0, "user_id": "", "team_ids": (), "api_key_hash": "hashed-key"}, "" - ) + receiver.get_span.assert_awaited_with("t1", "s1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "") def test_get_span_serves_ui_content_from_stored_payloads(client): @@ -284,9 +300,7 @@ def test_get_span_serves_ui_content_from_stored_payloads(client): def test_trace_detail_passes_scoped_reference(client, receiver): receiver.get_trace.return_value = TRACE_RESPONSE assert client.get("/v1/traces/t1?trace_ref=run-one").status_code == 200 - receiver.get_trace.assert_awaited_with( - "t1", {"all_teams": 0, "user_id": "", "team_ids": (), "api_key_hash": "hashed-key"}, "run-one" - ) + receiver.get_trace.assert_awaited_with("t1", {"all_teams": 0, "user_id": "user", "team_ids": ()}, "run-one") def test_invalid_export_and_cursor_are_client_errors(client, receiver): @@ -298,12 +312,34 @@ def test_invalid_export_and_cursor_are_client_errors(client, receiver): assert client.get("/v1/traces?cursor=broken").status_code == 400 -def test_teamless_key_without_token_gets_403_on_reads(client, receiver): - client.app.dependency_overrides[user_api_key_auth] = lambda: UserAPIKeyAuth( - user_role=LitellmUserRoles.INTERNAL_USER - ) - assert client.get("/v1/traces").status_code == 403 - receiver.list_traces.assert_not_called() +@pytest.mark.parametrize( + "auth", + ( + UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER), + UserAPIKeyAuth(token="key"), + UserAPIKeyAuth(token="key", team_id="unpermitted"), + UserAPIKeyAuth(user_id="", token="key"), + ), +) +def test_key_without_user_cannot_read_traces(client: TestClient, auth: UserAPIKeyAuth) -> None: + storage: Final = MagicMock(spec=ClickHouseStorage) + client.app.dependency_overrides[user_api_key_auth] = lambda: auth + client.app.dependency_overrides[tracing_endpoints.provide_receiver] = lambda: TraceReceiver(TraceStore(storage)) + client.app.dependency_overrides[tracing_endpoints.provide_trace_query_secret] = lambda: "test-secret" + for path in ( + "/v1/traces", + "/v1/traces/t1", + "/v1/traces/t1/spans/s1", + "/v1/traces/t1/spans/s1/error", + "/v1/traces/query/help", + ): + response: Final = client.get(path) + assert response.status_code == 403, response.text + query: Final = client.post("/v1/traces/query", json={"sql": "SELECT * FROM otel_traces"}) + assert query.status_code == 403, query.text + storage.query.assert_not_called() + storage.query_sql.assert_not_called() + storage.query_help.assert_not_called() def test_view_only_admin_cannot_ingest_traces(client, receiver): @@ -477,9 +513,8 @@ def test_lifespan_receivers_are_app_local() -> None: SPAN_DETAIL, SpanDetailParams( all_teams=0, - user_id="", + user_id=TEAM_KEY.user_id, team_ids=(), - api_key_hash=TEAM_KEY.token, trace_id="t1", span_id="first-span", trace_ref="first-run", @@ -489,9 +524,8 @@ def test_lifespan_receivers_are_app_local() -> None: SPAN_DETAIL, SpanDetailParams( all_teams=0, - user_id="", + user_id=TEAM_KEY.user_id, team_ids=(), - api_key_hash=TEAM_KEY.token, trace_id="t1", span_id="second-span", trace_ref="second-run", @@ -582,14 +616,17 @@ def test_lens_reads_from_injected_storage_without_receiver() -> None: @pytest.mark.parametrize( ("auth", "expected_scope"), ( - (UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), {"kind": "admin"}), - (UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY), {"kind": "admin"}), - (TEAM_KEY, {"kind": "logs", "user_id": "", "team_ids": (), "api_key_hash": "hashed-key"}), + (UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), {"kind": "all"}), + (UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY), {"kind": "all"}), + (TEAM_KEY, {"kind": "owned", "user_id": "user", "team_ids": ()}), ( - UserAPIKeyAuth(token="project-key", team_id="team-a", project_id="project-a"), - {"kind": "logs", "user_id": "", "team_ids": (), "api_key_hash": "project-key"}, + UserAPIKeyAuth(user_id="user", token="project-key", team_id="team-a", project_id="project-a"), + {"kind": "owned", "user_id": "user", "team_ids": ()}, + ), + ( + UserAPIKeyAuth(user_id="user", token="solo-key"), + {"kind": "owned", "user_id": "user", "team_ids": ()}, ), - (UserAPIKeyAuth(token="solo-key"), {"kind": "logs", "user_id": "", "team_ids": (), "api_key_hash": "solo-key"}), ), ) def test_sql_and_help_use_authenticated_scope( @@ -609,7 +646,7 @@ def test_sql_and_help_use_authenticated_scope( assert help_result.status_code == 200, help_result.text assert help_result.json() == QUERY_HELP receiver.store.storage.query_help.assert_awaited_once_with(expected_scope, "test-secret") - forged: Final = client.post("/v1/traces/query", json={"sql": "SELECT 1", "scope": {"kind": "admin"}}) + forged: Final = client.post("/v1/traces/query", json={"sql": "SELECT 1", "scope": {"kind": "all"}}) assert forged.status_code == 422, forged.text assert receiver.store.storage.query_sql.await_count == 1 @@ -638,7 +675,7 @@ def test_sql_reports_rejected_queries_and_unavailable_readers( result: Final = client.post("/v1/traces/query", json={"sql": "SELECT 1"}) assert result.status_code == status, result.text receiver.store.storage.query_sql.assert_awaited_once_with( - "SELECT 1", {"kind": "logs", "user_id": "", "team_ids": (), "api_key_hash": "hashed-key"}, "test-secret" + "SELECT 1", {"kind": "owned", "user_id": "user", "team_ids": ()}, "test-secret" ) @@ -648,7 +685,7 @@ def test_query_help_does_not_fall_back_when_reader_provisioning_fails(client: Te result: Final = client.get("/v1/traces/query/help") assert result.status_code == 503, result.text receiver.store.storage.query_help.assert_awaited_once_with( - {"kind": "logs", "user_id": "", "team_ids": (), "api_key_hash": "hashed-key"}, "test-secret" + {"kind": "owned", "user_id": "user", "team_ids": ()}, "test-secret" ) @@ -668,10 +705,96 @@ def test_queries_require_a_proxy_secret( return assert result.status_code == 200, result.text receiver.store.storage.query_sql.assert_awaited_once_with( - "SELECT 1", {"kind": "logs", "user_id": "", "team_ids": (), "api_key_hash": "hashed-key"}, secret + "SELECT 1", {"kind": "owned", "user_id": "user", "team_ids": ()}, secret ) +@pytest.mark.parametrize( + ("auth", "teams", "expected"), + ( + (UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), ("team-a",), (1, "", ())), + (UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY), ("team-a",), (1, "", ())), + (UserAPIKeyAuth(user_id="user", token="key", team_id="unpermitted"), ("a", "b"), (0, "user", ("a", "b"))), + (UserAPIKeyAuth(user_id="user", token="key"), (), (0, "user", ())), + (UserAPIKeyAuth(user_id="user"), ("a",), (0, "user", ("a",))), + ), +) +def test_shared_trace_permissions_reach_read_and_sql_boundaries( + client: TestClient, + auth: UserAPIKeyAuth, + teams: tuple[str, ...], + expected: tuple[Literal[0, 1], str, tuple[str, ...]], +) -> None: + async def lookup(caller: UserAPIKeyAuth) -> tuple[str, ...]: + assert caller is auth + return teams + + team_lookup: Final = AsyncMock(side_effect=lookup) + storage: Final = MagicMock(spec=ClickHouseStorage) + storage.query = AsyncMock(return_value=[{"span_id": "s1", "input": "", "output": "", "attributes": {}}]) + storage.query_sql = AsyncMock(return_value=TraceSQLResponse.model_validate(SQL_ENVELOPE)) + storage.query_help = AsyncMock(return_value=TraceQueryHelp.model_validate(QUERY_HELP)) + client.app.dependency_overrides[user_api_key_auth] = lambda: auth + client.app.dependency_overrides[get_log_team_lookup] = lambda: team_lookup + client.app.dependency_overrides[tracing_endpoints.provide_receiver] = lambda: TraceReceiver(TraceStore(storage)) + client.app.dependency_overrides[tracing_endpoints.provide_trace_query_secret] = lambda: "test-secret" + + response: Final = client.get("/v1/traces/t1/spans/s1?trace_ref=run-one") + assert response.status_code == 200, response.text + assert response.json()["span_id"] == "s1" + storage.query.assert_awaited_once_with( + SPAN_DETAIL, + SpanDetailParams( + all_teams=expected[0], + user_id=expected[1], + team_ids=expected[2], + trace_id="t1", + span_id="s1", + trace_ref="run-one", + ), + ) + sql_response: Final = client.post("/v1/traces/query", json={"sql": "SELECT * FROM otel_traces"}) + assert sql_response.status_code == 200, sql_response.text + assert sql_response.json() == SQL_ENVELOPE + assert client.get("/v1/traces/query/help").json() == QUERY_HELP + query_scope: Final = ( + {"kind": "all"} + if expected[0] + else { + "kind": "owned", + "user_id": expected[1], + "team_ids": expected[2], + } + ) + storage.query_sql.assert_awaited_once_with("SELECT * FROM otel_traces", query_scope, "test-secret") + storage.query_help.assert_awaited_once_with(query_scope, "test-secret") + assert team_lookup.await_count == ( + 3 + if auth.user_id and auth.user_role not in (LitellmUserRoles.PROXY_ADMIN, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY) + else 0 + ) + + +@pytest.mark.parametrize( + ("scope", "expected"), + ( + (OwnedRows(None), ("", ())), + (OwnedRows("user"), ("user", ())), + (OwnedRows("user", ("a", "b")), ("user", ("a", "b"))), + ), +) +def test_trace_storage_permissions_map_owned_rows( + scope: ReadScope, + expected: tuple[str, tuple[str, ...]], +) -> None: + assert tracing_endpoints._trace_scope(scope) == TraceScope(all_teams=0, user_id=expected[0], team_ids=expected[1]) + assert tracing_endpoints.trace_query_scope(scope) == { + "kind": "owned", + "user_id": expected[0], + "team_ids": expected[1], + } + + class _NativeConfig: def __init__(self, database: str, url: str, retention_days: int) -> None: pass @@ -685,7 +808,7 @@ class _NativeReturningHelp(ModuleType): def __init__(self, config: _NativeConfig) -> None: pass - async def query_help(self, scope: AdminQueryScope, secret: str) -> Mapping[str, object]: + async def query_help(self, scope: AllQueryScope, secret: str) -> Mapping[str, object]: return help_payload self.NativeTraceConfig: Final = _NativeConfig @@ -698,7 +821,7 @@ class _NativeReturningHelp(ModuleType): async def test_storage_validates_the_native_query_help_value(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(loader, "_cached_bridge", _NativeReturningHelp(QUERY_HELP)) storage: Final = ClickHouseStorage(TraceStorageConfig("http://clickhouse:8123")) - assert await storage.query_help({"kind": "admin"}, "secret") == TraceQueryHelp.model_validate(QUERY_HELP) + assert await storage.query_help({"kind": "all"}, "secret") == TraceQueryHelp.model_validate(QUERY_HELP) @pytest.mark.parametrize( @@ -726,4 +849,4 @@ async def test_storage_rejects_native_query_help_that_drifts_from_the_contract( monkeypatch.setattr(loader, "_cached_bridge", _NativeReturningHelp({**QUERY_HELP, **drift})) storage: Final = ClickHouseStorage(TraceStorageConfig("http://clickhouse:8123")) with pytest.raises(RuntimeError, match="invalid response"): - await storage.query_help({"kind": "admin"}, "secret") + await storage.query_help({"kind": "all"}, "secret") diff --git a/tests/unit/rust_bridge/test_trace_queries.py b/tests/unit/rust_bridge/test_trace_queries.py index 5ffe9ea0404..76ead3582a4 100644 --- a/tests/unit/rust_bridge/test_trace_queries.py +++ b/tests/unit/rust_bridge/test_trace_queries.py @@ -15,7 +15,6 @@ def test_named_query_rejects_offsets_outside_the_native_integer_range(offset: in "all_teams": 0, "user_id": "", "team_ids": ["team"], - "api_key_hash": "key", "trace_id": "trace", "trace_ref": "ref", "span_id": "span", @@ -31,7 +30,6 @@ def test_named_query_rejects_parameters_for_a_different_query() -> None: all_teams=0, user_id="", team_ids=("team",), - api_key_hash="key", trace_id="trace", trace_ref="ref", span_id="span", diff --git a/ui/litellm-dashboard/src/lib/http/schema.d.ts b/ui/litellm-dashboard/src/lib/http/schema.d.ts index df968e0bc9f..722f10cd571 100644 --- a/ui/litellm-dashboard/src/lib/http/schema.d.ts +++ b/ui/litellm-dashboard/src/lib/http/schema.d.ts @@ -46901,8 +46901,11 @@ export interface components { expression: string; /** Key */ key: string; - /** Type */ - type: string; + /** + * Type + * @constant + */ + type: "String"; }; /** TraceQueryAttributes */ TraceQueryAttributes: { @@ -46916,8 +46919,11 @@ export interface components { fields: components["schemas"]["TraceQueryAttributeField"][]; /** Scope */ scope: string; - /** Table */ - table: string; + /** + * Table + * @enum {string} + */ + table: "otel_traces" | "agent_traces_by_key" | "spend_logs"; /** Truncated */ truncated: boolean; }; @@ -46977,8 +46983,11 @@ export interface components { sampled_rows: number; /** Scope */ scope: string; - /** Table */ - table: string; + /** + * Table + * @enum {string} + */ + table: "otel_traces" | "agent_traces_by_key" | "spend_logs"; /** Truncated */ truncated: boolean; }; @@ -46989,7 +46998,7 @@ export interface components { /** Path */ path: (string | number)[]; /** Types */ - types: string[]; + types: ("array" | "boolean" | "integer" | "null" | "number" | "object" | "string")[]; }; /** TraceQueryNormalizedField */ TraceQueryNormalizedField: { @@ -46999,8 +47008,11 @@ export interface components { meaning: string; /** Name */ name: string; - /** Table */ - table: string; + /** + * Table + * @enum {string} + */ + table: "otel_traces" | "agent_traces_by_key" | "spend_logs"; /** Type */ type: string; }; @@ -47035,8 +47047,11 @@ export interface components { TraceQueryTable: { /** Columns */ columns: components["schemas"]["TraceQueryColumn"][]; - /** Name */ - name: string; + /** + * Name + * @enum {string} + */ + name: "otel_traces" | "agent_traces_by_key" | "spend_logs"; }; /** TraceSQLResponse */ TraceSQLResponse: {