From 04b1b077de07373613f3905f98f100f0b86aee04 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 07:54:54 -0700 Subject: [PATCH 01/10] fix(ui): keep untouched stored auto-router booleans and reminder marker casing on save (#42703) Co-authored-by: yuneng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ...d_updated_complexity_router_config.test.ts | 102 ++++++++++++++++++ .../edit_auto_router_modal.tsx | 18 +++- 2 files changed, 119 insertions(+), 1 deletion(-) diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts b/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts index 162414b7814..4b8e298191d 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts +++ b/ui/litellm-dashboard/src/components/edit_auto_router/build_updated_complexity_router_config.test.ts @@ -1115,3 +1115,105 @@ describe("LLM V2 configuration preservation", () => { expect(saved).not.toHaveProperty("classifier_llm_config"); }); }); + +describe("untouched save round trip", () => { + const STORED_PRE_MANAGED_BOOLEANS: Record = { + tiers: { SIMPLE: ["gpt-4o-mini"], MEDIUM: ["gpt-4o"], COMPLEX: ["opus"], REASONING: ["o1"] }, + tier_model_configs: { REASONING: [{ model_name: "o1", litellm_params: { reasoning_effort: "high" } }] }, + default_model: "gpt-4o", + plan_mode_min_tier: "COMPLEX", + tier_labels: { SIMPLE: "Cheap" }, + classifier_type: "heuristic_first", + heuristic_v2_success_threshold: 0.89, + heuristic_first_max_tier: "SIMPLE", + classifier_llm_config: { model: "gpt-4o-mini", timeout_ms: 3000, reasoning_effort: "low" }, + classifier_context_window_size: 5, + classifier_context_budget_chars: 4000, + classifier_context_include_assistant_turns: true, + classifier_fallback: "default_model", + classification_prompt: "Route for a payments team.", + classification_examples: "- refund status -> SIMPLE", + classification_mode: "user_turn", + session_affinity: true, + session_affinity_ttl_seconds: 300, + modality_routing: true, + modality_pin_override: true, + deployment_affinity: false, + adaptive: true, + adaptive_weights: { quality: 0.4, cost: 0.6 }, + tier_distance_penalty: 0.25, + adaptive_eligible: "all", + return_raw_model_name: true, + tier_boundaries: { simple_medium: 0.2, medium_complex: 0.4, complex_reasoning: 0.7 }, + token_thresholds: { simple: 20, complex: 500 }, + dimension_weights: { tokenCount: 0.1 }, + custom_dimensions: [{ name: "domain", weight: 0.9, keywords: ["orbitmesh"] }], + reasoning_override_min_score: 0.3, + enable_context_window_escalation: false, + context_window_escalation_buffer: 0.9, + code_keywords: ["async", "await"], + reasoning_keywords: ["prove"], + technical_keywords: ["api"], + simple_keywords: ["hello"], + plan_mode_patterns: ["plan now"], + route_housekeeping_to_cheapest_tier: true, + housekeeping_patterns: ["conversation title"], + reminder_markers: [{ open: "", close: "" }], + max_tokens_from_tier_model: true, + }; + + it("returns the stored config unchanged when nothing was edited", () => { + const hydrated = hydrateComplexityRouterConfig(STORED_PRE_MANAGED_BOOLEANS, undefined); + expect(buildUpdatedComplexityRouterConfig(STORED_PRE_MANAGED_BOOLEANS, hydrated)).toEqual( + STORED_PRE_MANAGED_BOOLEANS, + ); + }); + + it("keeps both housekeeping and max-token booleans stored as true", () => { + const hydrated = hydrateComplexityRouterConfig(STORED_PRE_MANAGED_BOOLEANS, undefined); + const saved = buildUpdatedComplexityRouterConfig(STORED_PRE_MANAGED_BOOLEANS, hydrated); + expect(saved.route_housekeeping_to_cheapest_tier).toBe(true); + expect(saved.max_tokens_from_tier_model).toBe(true); + }); + + it("keeps the stored reminder marker casing", () => { + const hydrated = hydrateComplexityRouterConfig(STORED_PRE_MANAGED_BOOLEANS, undefined); + const saved = buildUpdatedComplexityRouterConfig(STORED_PRE_MANAGED_BOOLEANS, hydrated); + expect(saved.reminder_markers).toEqual(STORED_PRE_MANAGED_BOOLEANS.reminder_markers); + }); + + it("lets an edited toggle win over the stored value", () => { + const hydrated = hydrateComplexityRouterConfig(STORED_PRE_MANAGED_BOOLEANS, undefined); + const edited = { + ...hydrated, + route_housekeeping_to_cheapest_tier: false, + max_tokens_from_tier_model: false, + }; + const saved = buildUpdatedComplexityRouterConfig(STORED_PRE_MANAGED_BOOLEANS, edited); + expect(saved.route_housekeeping_to_cheapest_tier).toBe(false); + expect(saved.max_tokens_from_tier_model).toBe(false); + + const storedDisabled: Record = { + ...STORED_PRE_MANAGED_BOOLEANS, + route_housekeeping_to_cheapest_tier: false, + max_tokens_from_tier_model: false, + }; + const enabled = { + ...hydrateComplexityRouterConfig(storedDisabled, undefined), + route_housekeeping_to_cheapest_tier: true, + max_tokens_from_tier_model: true, + }; + const resaved = buildUpdatedComplexityRouterConfig(storedDisabled, enabled); + expect(resaved).not.toHaveProperty("route_housekeeping_to_cheapest_tier"); + expect(resaved).not.toHaveProperty("max_tokens_from_tier_model"); + }); + + it("lowercases reminder markers the user edited", () => { + const hydrated = hydrateComplexityRouterConfig(STORED_PRE_MANAGED_BOOLEANS, undefined); + const saved = buildUpdatedComplexityRouterConfig(STORED_PRE_MANAGED_BOOLEANS, { + ...hydrated, + reminder_markers: [{ open: "", close: "" }], + }); + expect(saved.reminder_markers).toEqual([{ open: "", close: "" }]); + }); +}); diff --git a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx index 0429e7d4ad1..2fd4fba2b31 100644 --- a/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx +++ b/ui/litellm-dashboard/src/components/edit_auto_router/edit_auto_router_modal.tsx @@ -33,6 +33,7 @@ import { import { isComplexityRouter } from "../add_model/auto_router_strategies"; import { type BuildComplexityRouterConfigParams, + type StoredComplexityRouterConfig, buildComplexityRouterConfig, getClassifierModelError, getHeuristicV2SuccessThresholdError, @@ -152,6 +153,12 @@ const KEYWORD_MATCHING_KEYS = new Set([ "match_threshold", ]); +const UNEDITED_STORED_VALUE_KEYS: readonly (keyof ComplexityRouterConfigValue)[] = [ + "route_housekeeping_to_cheapest_tier", + "reminder_markers", + "max_tokens_from_tier_model", +]; + const toRecord = (value: unknown): Record => { const parsed: unknown = typeof value === "string" ? JSON.parse(value) : value; return typeof parsed === "object" && parsed !== null && !Array.isArray(parsed) @@ -179,7 +186,15 @@ export const buildUpdatedComplexityRouterConfig = ( customTechnicalKeywords?: string[], keywordMatching?: KeywordMatchingState, ): Record => { + const stored = toRecord(storedConfig); + const hydratedFromStored = hydrateComplexityRouterConfig(stored as StoredComplexityRouterConfig, undefined); + const unedited = new Set( + UNEDITED_STORED_VALUE_KEYS.filter( + (key) => key in stored && JSON.stringify(value[key]) === JSON.stringify(hydratedFromStored[key]), + ), + ); const isManaged = (key: string): boolean => { + if (unedited.has(key)) return false; if (key === "classifier_context_per_turn_chars") { return !usesClassifierContext(effectiveClassifierType(value)) || Object.prototype.hasOwnProperty.call(value, key); } @@ -190,7 +205,7 @@ export const buildUpdatedComplexityRouterConfig = ( }; const dropped = customTierDroppedKeys(value); const preservedConfig = Object.fromEntries( - Object.entries(toRecord(storedConfig)).filter(([key]) => !isManaged(key) && !dropped.includes(key)), + Object.entries(stored).filter(([key]) => !isManaged(key) && !dropped.includes(key)), ); const builderParams: BuildComplexityRouterConfigParams = { @@ -208,6 +223,7 @@ export const buildUpdatedComplexityRouterConfig = ( const unowned: readonly string[] = [ ...(keywordMatching === undefined ? [...KEYWORD_MATCHING_KEYS].filter((key) => !isManaged(key)) : []), ...(customTechnicalKeywords === undefined ? ["custom_technical_keywords"] : []), + ...unedited, ]; return { ...preservedConfig, From 75a6bca8b90e2a477004bf30d66e9978d0efb30e Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 08:13:41 -0700 Subject: [PATCH 02/10] feat(bedrock): add gpt-6-sol and gpt-6-luna model pricing (#42746) Co-authored-by: kerry Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 204 ++++++++++++++++++ model_prices_and_context_window.json | 204 ++++++++++++++++++ 2 files changed, 408 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index b8e4331dfa2..98bda7c12c8 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -55684,6 +55684,82 @@ "supports_vision": true, "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-astra.html" }, + "bedrock_mantle/openai.gpt-6-sol": { + "input_cost_per_token": 2.2e-06, + "input_cost_per_token_above_272k_tokens": 4.4e-06, + "cache_creation_input_token_cost": 2.75e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5.5e-06, + "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost_above_272k_tokens": 4.4e-07, + "output_cost_per_token": 1.1e-05, + "output_cost_per_token_above_272k_tokens": 1.65e-05, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "responses", + "use_openai_responses_path": true, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, + "supports_vision": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-cards-openai.html" + }, + "bedrock_mantle/openai.gpt-6-luna": { + "input_cost_per_token": 1.1e-07, + "input_cost_per_token_above_272k_tokens": 2.2e-07, + "cache_creation_input_token_cost": 1.375e-07, + "cache_creation_input_token_cost_above_272k_tokens": 2.75e-07, + "cache_read_input_token_cost": 1.1e-08, + "cache_read_input_token_cost_above_272k_tokens": 2.2e-08, + "output_cost_per_token": 5.5e-07, + "output_cost_per_token_above_272k_tokens": 8.25e-07, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "responses", + "use_openai_responses_path": true, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, + "supports_vision": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-cards-openai.html" + }, "us.openai.gpt-6-astra": { "input_cost_per_token": 1.1e-05, "input_cost_per_token_above_272k_tokens": 2.2e-05, @@ -55716,6 +55792,70 @@ "supports_vision": true, "source": "https://aws.amazon.com/bedrock/pricing/" }, + "us.openai.gpt-6-sol": { + "input_cost_per_token": 2.2e-06, + "input_cost_per_token_above_272k_tokens": 4.4e-06, + "cache_creation_input_token_cost": 2.75e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5.5e-06, + "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost_above_272k_tokens": 4.4e-07, + "output_cost_per_token": 1.1e-05, + "output_cost_per_token_above_272k_tokens": 1.65e-05, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_tool_choice": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, + "supports_vision": true, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, + "us.openai.gpt-6-luna": { + "input_cost_per_token": 1.1e-07, + "input_cost_per_token_above_272k_tokens": 2.2e-07, + "cache_creation_input_token_cost": 1.375e-07, + "cache_creation_input_token_cost_above_272k_tokens": 2.75e-07, + "cache_read_input_token_cost": 1.1e-08, + "cache_read_input_token_cost_above_272k_tokens": 2.2e-08, + "output_cost_per_token": 5.5e-07, + "output_cost_per_token_above_272k_tokens": 8.25e-07, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_tool_choice": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, + "supports_vision": true, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, "global.openai.gpt-6-astra": { "input_cost_per_token": 1e-05, "input_cost_per_token_above_272k_tokens": 2e-05, @@ -55748,6 +55888,70 @@ "supports_vision": true, "source": "https://aws.amazon.com/bedrock/pricing/" }, + "global.openai.gpt-6-sol": { + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_tool_choice": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, + "supports_vision": true, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, + "global.openai.gpt-6-luna": { + "input_cost_per_token": 1e-07, + "input_cost_per_token_above_272k_tokens": 2e-07, + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_272k_tokens": 2.5e-07, + "cache_read_input_token_cost": 1e-08, + "cache_read_input_token_cost_above_272k_tokens": 2e-08, + "output_cost_per_token": 5e-07, + "output_cost_per_token_above_272k_tokens": 7.5e-07, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_tool_choice": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, + "supports_vision": true, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, "bedrock_mantle/openai.gpt-5.5": { "input_cost_per_token": 5.5e-06, "input_cost_per_token_above_272k_tokens": 1.1e-05, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index b8e4331dfa2..98bda7c12c8 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -55684,6 +55684,82 @@ "supports_vision": true, "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-card-openai-gpt-6-astra.html" }, + "bedrock_mantle/openai.gpt-6-sol": { + "input_cost_per_token": 2.2e-06, + "input_cost_per_token_above_272k_tokens": 4.4e-06, + "cache_creation_input_token_cost": 2.75e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5.5e-06, + "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost_above_272k_tokens": 4.4e-07, + "output_cost_per_token": 1.1e-05, + "output_cost_per_token_above_272k_tokens": 1.65e-05, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "responses", + "use_openai_responses_path": true, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, + "supports_vision": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-cards-openai.html" + }, + "bedrock_mantle/openai.gpt-6-luna": { + "input_cost_per_token": 1.1e-07, + "input_cost_per_token_above_272k_tokens": 2.2e-07, + "cache_creation_input_token_cost": 1.375e-07, + "cache_creation_input_token_cost_above_272k_tokens": 2.75e-07, + "cache_read_input_token_cost": 1.1e-08, + "cache_read_input_token_cost_above_272k_tokens": 2.2e-08, + "output_cost_per_token": 5.5e-07, + "output_cost_per_token_above_272k_tokens": 8.25e-07, + "litellm_provider": "bedrock_mantle", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "responses", + "use_openai_responses_path": true, + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_tool_choice": true, + "supports_xhigh_reasoning_effort": true, + "supports_vision": true, + "source": "https://docs.aws.amazon.com/bedrock/latest/userguide/model-cards-openai.html" + }, "us.openai.gpt-6-astra": { "input_cost_per_token": 1.1e-05, "input_cost_per_token_above_272k_tokens": 2.2e-05, @@ -55716,6 +55792,70 @@ "supports_vision": true, "source": "https://aws.amazon.com/bedrock/pricing/" }, + "us.openai.gpt-6-sol": { + "input_cost_per_token": 2.2e-06, + "input_cost_per_token_above_272k_tokens": 4.4e-06, + "cache_creation_input_token_cost": 2.75e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5.5e-06, + "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost_above_272k_tokens": 4.4e-07, + "output_cost_per_token": 1.1e-05, + "output_cost_per_token_above_272k_tokens": 1.65e-05, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_tool_choice": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, + "supports_vision": true, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, + "us.openai.gpt-6-luna": { + "input_cost_per_token": 1.1e-07, + "input_cost_per_token_above_272k_tokens": 2.2e-07, + "cache_creation_input_token_cost": 1.375e-07, + "cache_creation_input_token_cost_above_272k_tokens": 2.75e-07, + "cache_read_input_token_cost": 1.1e-08, + "cache_read_input_token_cost_above_272k_tokens": 2.2e-08, + "output_cost_per_token": 5.5e-07, + "output_cost_per_token_above_272k_tokens": 8.25e-07, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_tool_choice": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, + "supports_vision": true, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, "global.openai.gpt-6-astra": { "input_cost_per_token": 1e-05, "input_cost_per_token_above_272k_tokens": 2e-05, @@ -55748,6 +55888,70 @@ "supports_vision": true, "source": "https://aws.amazon.com/bedrock/pricing/" }, + "global.openai.gpt-6-sol": { + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_tool_choice": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, + "supports_vision": true, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, + "global.openai.gpt-6-luna": { + "input_cost_per_token": 1e-07, + "input_cost_per_token_above_272k_tokens": 2e-07, + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_272k_tokens": 2.5e-07, + "cache_read_input_token_cost": 1e-08, + "cache_read_input_token_cost_above_272k_tokens": 2e-08, + "output_cost_per_token": 5e-07, + "output_cost_per_token_above_272k_tokens": 7.5e-07, + "litellm_provider": "bedrock_converse", + "max_input_tokens": 1050000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_none_reasoning_effort": false, + "supports_tool_choice": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_xhigh_reasoning_effort": true, + "supports_vision": true, + "source": "https://aws.amazon.com/bedrock/pricing/" + }, "bedrock_mantle/openai.gpt-5.5": { "input_cost_per_token": 5.5e-06, "input_cost_per_token_above_272k_tokens": 1.1e-05, From 48050d9646d0edd99b9e340b694918e9ee5b444e Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 08:14:11 -0700 Subject: [PATCH 03/10] chore(ci): drop litellm_internal_staging and litellm_oss_staging references, main is the only trunk (#42745) --- .circleci/config.yml | 2 +- .circleci/scripts/path_filter.sh | 2 +- .githooks/pre-push | 6 +-- .github/workflows/check-ui-api-types.yml | 2 - .github/workflows/ci-coverage.yml | 3 -- .github/workflows/codspeed.yml | 2 - .github/workflows/cost-map-guard.yml | 2 - .../workflows/create_daily_staging_branch.yml | 49 ------------------- .github/workflows/guard-fork-dependencies.yml | 2 - .github/workflows/image-scan.yml | 1 - .github/workflows/osv-scan.yml | 2 - .../publish-basedpyright-base-counts.yml | 3 +- .github/workflows/test-code-quality.yml | 3 -- .github/workflows/test-linting.yml | 2 - .github/workflows/test-litellm-ui-build.yml | 2 - .github/workflows/test-litellm-ui-lint.yml | 2 - .github/workflows/test-litellm-ui-unit.yml | 3 -- .../test-mcp-dependency-resolution.yml | 2 - .github/workflows/test-postgres.yml | 3 -- .github/workflows/test-redis-compat.yml | 2 - .github/workflows/test-rust.yml | 2 - .github/workflows/test-semgrep.yml | 2 - .github/workflows/test-terraform-modules.yml | 2 - .github/workflows/test-terraform-provider.yml | 2 - .github/workflows/test-unit-documentation.yml | 3 -- .github/workflows/test-unit-proxy-db.yml | 3 -- .github/workflows/test-unit.yml | 3 -- .github/workflows/test-vscode-extension.yml | 2 - .github/workflows/zizmor.yml | 4 +- scripts/comment-fixed-issue.test.ts | 22 ++++----- scripts/type_check_gate.py | 2 +- tests/test_litellm/test_default_branch.py | 18 +++---- tests/test_litellm/test_git_hooks.py | 1 - tests/test_litellm/test_pre_commit_lint.py | 4 +- .../test_litellm/test_select_ui_test_scope.py | 2 +- 35 files changed, 30 insertions(+), 137 deletions(-) delete mode 100644 .github/workflows/create_daily_staging_branch.yml diff --git a/.circleci/config.yml b/.circleci/config.yml index 7241502b12c..eb76244c1ab 100644 --- a/.circleci/config.yml +++ b/.circleci/config.yml @@ -3222,7 +3222,7 @@ workflows: cron: "17 0,6,12,18 * * *" filters: branches: - only: litellm_internal_staging + only: main jobs: *migration_jobs integration: unless: << pipeline.parameters.run_migration_tests >> diff --git a/.circleci/scripts/path_filter.sh b/.circleci/scripts/path_filter.sh index cdadde732bd..1da29f99f6a 100755 --- a/.circleci/scripts/path_filter.sh +++ b/.circleci/scripts/path_filter.sh @@ -11,7 +11,7 @@ run_full() { [ -n "${CIRCLE_PULL_REQUEST:-}" ] || run_full "not a pull request" -candidate_bases="main litellm_internal_staging litellm_oss_staging" +candidate_bases="main" merge_base="" for base in $candidate_bases; do git fetch --quiet origin "$base" 2>/dev/null || continue diff --git a/.githooks/pre-push b/.githooks/pre-push index c2267c8501c..dc1a73a7ba2 100755 --- a/.githooks/pre-push +++ b/.githooks/pre-push @@ -8,7 +8,6 @@ # # Protected branches (always allowed): # - main -# - litellm_internal_staging # - dependabot/* # - gh-readonly-queue/* # @@ -22,7 +21,7 @@ ZERO_OID_SHA256="000000000000000000000000000000000000000000000000000000000000000 ALLOWED_TYPES="feature|bugfix|hotfix|release|chore" BRANCH_PATTERN="^(${ALLOWED_TYPES})/.+" -PROTECTED_NAMES="main litellm_internal_staging" +PROTECTED_NAMES="main" PROTECTED_PREFIXES="dependabot/ gh-readonly-queue/" is_protected() { @@ -78,8 +77,7 @@ if [ -n "$invalid" ]; then chore/bump-deps hotfix/auth-bypass - Protected (always allowed): main, litellm_internal_staging, - dependabot/*, gh-readonly-queue/*. + Protected (always allowed): main, dependabot/*, gh-readonly-queue/*. See https://conventional-branch.github.io/ diff --git a/.github/workflows/check-ui-api-types.yml b/.github/workflows/check-ui-api-types.yml index 312a80103f8..185c20d916d 100644 --- a/.github/workflows/check-ui-api-types.yml +++ b/.github/workflows/check-ui-api-types.yml @@ -4,8 +4,6 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" permissions: diff --git a/.github/workflows/ci-coverage.yml b/.github/workflows/ci-coverage.yml index 7bc476db134..cd36a9a7ed6 100644 --- a/.github/workflows/ci-coverage.yml +++ b/.github/workflows/ci-coverage.yml @@ -4,13 +4,10 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" push: branches: - main - - litellm_internal_staging permissions: contents: read diff --git a/.github/workflows/codspeed.yml b/.github/workflows/codspeed.yml index ec7e211faa1..dfad57dc5ab 100644 --- a/.github/workflows/codspeed.yml +++ b/.github/workflows/codspeed.yml @@ -4,7 +4,6 @@ on: push: branches: - main - - litellm_internal_staging paths: - "litellm/**" - "tests/benchmarks/**" @@ -17,7 +16,6 @@ on: pull_request: branches: - main - - litellm_internal_staging paths: - "litellm/**" - "tests/benchmarks/**" diff --git a/.github/workflows/cost-map-guard.yml b/.github/workflows/cost-map-guard.yml index 61a56f1f47e..208a75434a5 100644 --- a/.github/workflows/cost-map-guard.yml +++ b/.github/workflows/cost-map-guard.yml @@ -4,8 +4,6 @@ on: # zizmor: ignore[dangerous-triggers] runs the base branch's code only; the P pull_request_target: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" permissions: diff --git a/.github/workflows/create_daily_staging_branch.yml b/.github/workflows/create_daily_staging_branch.yml deleted file mode 100644 index 6422b0d4dbc..00000000000 --- a/.github/workflows/create_daily_staging_branch.yml +++ /dev/null @@ -1,49 +0,0 @@ -name: Create Daily Staging Branch - -on: - schedule: - - cron: "0 0,12 * * *" # Runs every 12 hours at midnight and noon UTC - workflow_dispatch: # Allow manual trigger - -jobs: - create-staging-branch: - if: github.repository == 'BerriAI/litellm' - runs-on: ubuntu-latest - permissions: - contents: write - - steps: - - name: Create daily staging branch - env: - GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} - run: | - BRANCH_NAME="litellm_oss_staging_$(date +'%m_%d_%Y')" - echo "Creating branch: $BRANCH_NAME" - if gh api "repos/${{ github.repository }}/git/ref/heads/$BRANCH_NAME" --silent 2>/dev/null; then - echo "Branch $BRANCH_NAME already exists. Skipping creation." - exit 0 - fi - MAIN_SHA=$(gh api "repos/${{ github.repository }}/git/ref/heads/main" --jq '.object.sha') - gh api "repos/${{ github.repository }}/git/refs" -f ref="refs/heads/$BRANCH_NAME" -f sha="$MAIN_SHA" --silent - echo "Successfully created branch: $BRANCH_NAME at $MAIN_SHA" - - create-internal-dev-branch: - if: github.repository == 'BerriAI/litellm' - runs-on: ubuntu-latest - permissions: - contents: write - - steps: - - name: Create internal dev branch - env: - GH_TOKEN: ${{ secrets.GITHUB_TOKEN }} - run: | - BRANCH_NAME="litellm_internal_dev_$(date +'%m_%d_%Y')" - echo "Creating branch: $BRANCH_NAME" - if gh api "repos/${{ github.repository }}/git/ref/heads/$BRANCH_NAME" --silent 2>/dev/null; then - echo "Branch $BRANCH_NAME already exists. Skipping creation." - exit 0 - fi - MAIN_SHA=$(gh api "repos/${{ github.repository }}/git/ref/heads/main" --jq '.object.sha') - gh api "repos/${{ github.repository }}/git/refs" -f ref="refs/heads/$BRANCH_NAME" -f sha="$MAIN_SHA" --silent - echo "Successfully created branch: $BRANCH_NAME at $MAIN_SHA" diff --git a/.github/workflows/guard-fork-dependencies.yml b/.github/workflows/guard-fork-dependencies.yml index 6b366da78d4..a0717c8e0f9 100644 --- a/.github/workflows/guard-fork-dependencies.yml +++ b/.github/workflows/guard-fork-dependencies.yml @@ -4,8 +4,6 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" paths: - "uv.lock" diff --git a/.github/workflows/image-scan.yml b/.github/workflows/image-scan.yml index c27d49ed610..0695720733f 100644 --- a/.github/workflows/image-scan.yml +++ b/.github/workflows/image-scan.yml @@ -4,7 +4,6 @@ on: pull_request: branches: - main - - litellm_internal_staging - litellm_oss_branch - "litellm_**" paths: diff --git a/.github/workflows/osv-scan.yml b/.github/workflows/osv-scan.yml index 0aedeaec12a..1c7e8b2841a 100644 --- a/.github/workflows/osv-scan.yml +++ b/.github/workflows/osv-scan.yml @@ -4,8 +4,6 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" schedule: - cron: "23 6 * * *" diff --git a/.github/workflows/publish-basedpyright-base-counts.yml b/.github/workflows/publish-basedpyright-base-counts.yml index 27d4682dbd9..34f60d25980 100644 --- a/.github/workflows/publish-basedpyright-base-counts.yml +++ b/.github/workflows/publish-basedpyright-base-counts.yml @@ -1,6 +1,6 @@ name: Publish basedpyright base counts -# Every commit on main or litellm_internal_staging can become a future merge-base. +# Every commit on main can become a future merge-base. # Publishing its per-rule basedpyright counts as an artifact lets # scripts/type_check_gate.py download them in seconds instead of paying a # 60-110s second basedpyright pass on every fresh worktree or moved merge-base. @@ -11,7 +11,6 @@ on: push: branches: - main - - litellm_internal_staging workflow_dispatch: inputs: ref: diff --git a/.github/workflows/test-code-quality.yml b/.github/workflows/test-code-quality.yml index 987f66773f2..2b6adaa6af6 100644 --- a/.github/workflows/test-code-quality.yml +++ b/.github/workflows/test-code-quality.yml @@ -4,13 +4,10 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" push: branches: - main - - litellm_internal_staging permissions: contents: read diff --git a/.github/workflows/test-linting.yml b/.github/workflows/test-linting.yml index 592d8edf6b8..c77d4b2ee96 100644 --- a/.github/workflows/test-linting.yml +++ b/.github/workflows/test-linting.yml @@ -4,8 +4,6 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" permissions: diff --git a/.github/workflows/test-litellm-ui-build.yml b/.github/workflows/test-litellm-ui-build.yml index 4eb6b272c43..eace78fc2cb 100644 --- a/.github/workflows/test-litellm-ui-build.yml +++ b/.github/workflows/test-litellm-ui-build.yml @@ -7,8 +7,6 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" concurrency: diff --git a/.github/workflows/test-litellm-ui-lint.yml b/.github/workflows/test-litellm-ui-lint.yml index e03d89ee26a..9ea5100e21b 100644 --- a/.github/workflows/test-litellm-ui-lint.yml +++ b/.github/workflows/test-litellm-ui-lint.yml @@ -6,8 +6,6 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" concurrency: diff --git a/.github/workflows/test-litellm-ui-unit.yml b/.github/workflows/test-litellm-ui-unit.yml index b93bf84320d..ee1440c6e8b 100644 --- a/.github/workflows/test-litellm-ui-unit.yml +++ b/.github/workflows/test-litellm-ui-unit.yml @@ -7,13 +7,10 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" push: branches: - main - - litellm_internal_staging concurrency: group: ${{ github.workflow }}-${{ github.event.pull_request.number || github.ref }} diff --git a/.github/workflows/test-mcp-dependency-resolution.yml b/.github/workflows/test-mcp-dependency-resolution.yml index b5dc573c1c1..463d6a7e9e8 100644 --- a/.github/workflows/test-mcp-dependency-resolution.yml +++ b/.github/workflows/test-mcp-dependency-resolution.yml @@ -4,8 +4,6 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" permissions: diff --git a/.github/workflows/test-postgres.yml b/.github/workflows/test-postgres.yml index 96c514dff7c..a1e6bf54135 100644 --- a/.github/workflows/test-postgres.yml +++ b/.github/workflows/test-postgres.yml @@ -4,13 +4,10 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" push: branches: - main - - litellm_internal_staging workflow_dispatch: permissions: diff --git a/.github/workflows/test-redis-compat.yml b/.github/workflows/test-redis-compat.yml index f29755a74b1..7862481173b 100644 --- a/.github/workflows/test-redis-compat.yml +++ b/.github/workflows/test-redis-compat.yml @@ -4,8 +4,6 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" paths: - "litellm/_redis.py" diff --git a/.github/workflows/test-rust.yml b/.github/workflows/test-rust.yml index 64e16b512e1..1f3b5c4d97c 100644 --- a/.github/workflows/test-rust.yml +++ b/.github/workflows/test-rust.yml @@ -29,8 +29,6 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" paths: - "litellm-rust/**" diff --git a/.github/workflows/test-semgrep.yml b/.github/workflows/test-semgrep.yml index 6e9f5e42fa2..d375824e3c9 100644 --- a/.github/workflows/test-semgrep.yml +++ b/.github/workflows/test-semgrep.yml @@ -4,8 +4,6 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" permissions: diff --git a/.github/workflows/test-terraform-modules.yml b/.github/workflows/test-terraform-modules.yml index e6896604b7f..e2fda3207e8 100644 --- a/.github/workflows/test-terraform-modules.yml +++ b/.github/workflows/test-terraform-modules.yml @@ -9,8 +9,6 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" paths: - "terraform/litellm/aws/**" diff --git a/.github/workflows/test-terraform-provider.yml b/.github/workflows/test-terraform-provider.yml index eb7b299fd1f..be7fd1e61dc 100644 --- a/.github/workflows/test-terraform-provider.yml +++ b/.github/workflows/test-terraform-provider.yml @@ -8,8 +8,6 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" paths: - "terraform/provider/**" diff --git a/.github/workflows/test-unit-documentation.yml b/.github/workflows/test-unit-documentation.yml index 90b6b28374e..660c7689e2b 100644 --- a/.github/workflows/test-unit-documentation.yml +++ b/.github/workflows/test-unit-documentation.yml @@ -4,13 +4,10 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" push: branches: - main - - litellm_internal_staging permissions: contents: read diff --git a/.github/workflows/test-unit-proxy-db.yml b/.github/workflows/test-unit-proxy-db.yml index 9013f21931b..4af7a161984 100644 --- a/.github/workflows/test-unit-proxy-db.yml +++ b/.github/workflows/test-unit-proxy-db.yml @@ -4,13 +4,10 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" push: branches: - main - - litellm_internal_staging permissions: contents: read diff --git a/.github/workflows/test-unit.yml b/.github/workflows/test-unit.yml index ac60a9b5f05..ed87049f2d5 100644 --- a/.github/workflows/test-unit.yml +++ b/.github/workflows/test-unit.yml @@ -4,13 +4,10 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" push: branches: - main - - litellm_internal_staging workflow_dispatch: permissions: diff --git a/.github/workflows/test-vscode-extension.yml b/.github/workflows/test-vscode-extension.yml index 886268d9e2c..c807948e4e9 100644 --- a/.github/workflows/test-vscode-extension.yml +++ b/.github/workflows/test-vscode-extension.yml @@ -6,8 +6,6 @@ on: pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" paths: - "vscode-extension/**" diff --git a/.github/workflows/zizmor.yml b/.github/workflows/zizmor.yml index df242e5a3b6..73f9efb8df9 100644 --- a/.github/workflows/zizmor.yml +++ b/.github/workflows/zizmor.yml @@ -2,12 +2,10 @@ name: GitHub Actions Security Analysis on: push: - branches: [main, litellm_internal_staging] + branches: [main] pull_request: branches: - main - - litellm_internal_staging - - litellm_oss_staging - "litellm_**" concurrency: diff --git a/scripts/comment-fixed-issue.test.ts b/scripts/comment-fixed-issue.test.ts index fef5c0f4007..f4242204938 100644 --- a/scripts/comment-fixed-issue.test.ts +++ b/scripts/comment-fixed-issue.test.ts @@ -263,7 +263,7 @@ describe("fixedBody", () => { describe("fixOf", () => { test("a merged pull request or a commit is a fix, whatever branch it was merged into", () => { expect(fixOf(closure(mergedPr), "BerriAI/litellm")).toEqual({ kind: "pull_request", number: 41767, oid: MERGE_COMMIT }); - expect(fixOf(closure({ ...mergedPr, baseRefName: "litellm_internal_staging" }), "BerriAI/litellm")).toEqual({ kind: "pull_request", number: 41767, oid: MERGE_COMMIT }); + expect(fixOf(closure({ ...mergedPr, baseRefName: "release_branch" }), "BerriAI/litellm")).toEqual({ kind: "pull_request", number: 41767, oid: MERGE_COMMIT }); expect(fixOf(closure(commitCloser), "BerriAI/litellm")).toEqual({ kind: "commit", oid: MERGE_COMMIT }); }); @@ -315,7 +315,7 @@ describe("closeVerdict", () => { for (const base of ["release/v1.102.0-rc.2", "stable/v1.83.14", "v_1_83_3_stable_patch"]) { expect(closeVerdict(openPr(41760, { baseRefName: base }), config)).toEqual({ kind: "skip", reason: `targets the release line ${base}` }); } - expect(closeVerdict(openPr(41760, { baseRefName: "litellm_internal_staging" }), config).kind).toBe("candidate"); + expect(closeVerdict(openPr(41760, { baseRefName: "release_branch" }), config).kind).toBe("candidate"); }); test("a pull request in a fork, one that is not open, or one linking no issue is left alone", () => { @@ -404,8 +404,8 @@ describe("handleFixedIssue", () => { }); test("a fix merged into the retired development branch or pushed as a commit counts once it is on the default branch", async () => { - const stagingPr = { ...mergedPr, baseRefName: "litellm_internal_staging" }; - const staging = openPr(41760, { baseRefName: "litellm_internal_staging", closingIssuesReferences: links(linkedIssue(ISSUE, stagingPr)) }); + const stagingPr = { ...mergedPr, baseRefName: "release_branch" }; + const staging = openPr(41760, { baseRefName: "release_branch", closingIssuesReferences: links(linkedIssue(ISSUE, stagingPr)) }); const byCommit = openPr(41761, { closingIssuesReferences: links(linkedIssue(41751, commitCloser)) }); const { api, writes } = fakeApi({ issue: closedBy(mergedPr, "CLOSED", [staging, byCommit]) }); const { pullRequests } = await handleFixedIssue(api, config, ISSUE, noPause); @@ -426,15 +426,15 @@ describe("handleFixedIssue", () => { }); test("the default branch comes from the config for the containment check and the comment alike", async () => { - const stagingConfig = { ...config, defaultBranch: "litellm_internal_staging" }; - const closer = { ...mergedPr, baseRefName: "litellm_internal_staging" }; + const stagingConfig = { ...config, defaultBranch: "release_branch" }; + const closer = { ...mergedPr, baseRefName: "release_branch" }; const { api, writes } = fakeApi({ issue: closedBy(closer, "CLOSED", [openPr(41760, { closingIssuesReferences: links(linkedIssue(ISSUE, closer)) })]), - reachable: { litellm_internal_staging: [MERGE_COMMIT] }, + reachable: { release_branch: [MERGE_COMMIT] }, }); const { pullRequests } = await handleFixedIssue(api, stagingConfig, ISSUE, noPause); - expect(pullRequests).toEqual([{ kind: "closed", number: 41760, body: supersededBody([prFix()], "litellm_internal_staging") }]); - expect(writes[1]).toContain("on litellm_internal_staging, so this pull request is closed"); + expect(pullRequests).toEqual([{ kind: "closed", number: 41760, body: supersededBody([prFix()], "release_branch") }]); + expect(writes[1]).toContain("on release_branch, so this pull request is closed"); }); test("the closer sits in the linked list as merged and gets neither a line nor a write", async () => { @@ -533,11 +533,11 @@ describe("handleFixedIssue", () => { }); test("an issue closed from the retired development branch gets no comment but still closes its linked pull requests once the fix is on the default branch", async () => { - const stagingPr = { ...mergedPr, baseRefName: "litellm_internal_staging" }; + const stagingPr = { ...mergedPr, baseRefName: "release_branch" }; const staging = openPr(41760, { closingIssuesReferences: links(linkedIssue(ISSUE, stagingPr)) }); const { api, writes } = fakeApi({ issue: closedBy(stagingPr, "CLOSED", [staging]) }); const { comment, pullRequests } = await handleFixedIssue(api, config, ISSUE, noPause); - expect(comment).toEqual({ kind: "skip", reason: "#41767 merged into litellm_internal_staging, not main" }); + expect(comment).toEqual({ kind: "skip", reason: "#41767 merged into release_branch, not main" }); expect(pullRequests).toEqual([{ kind: "closed", number: 41760, body: oneFixBody }]); expect(writes.map((write) => write.split(" ")[1])).toEqual(["/repos/BerriAI/litellm/issues/41760/comments", "/repos/BerriAI/litellm/pulls/41760"]); }); diff --git a/scripts/type_check_gate.py b/scripts/type_check_gate.py index 78f74ec65a1..023b35af8fe 100644 --- a/scripts/type_check_gate.py +++ b/scripts/type_check_gate.py @@ -35,7 +35,7 @@ detached worktree at the merge-base, run under the same environment so import resolution matches, and its per-rule counts are cached under the repo's git common dir keyed by merge-base commit, ``pyrightconfig.json``, ``uv.lock``, the Prisma schema, and the dependency-group set, so re-runs against the same -branch point pay for it once. A CI workflow publishes every staging commit's counts as +branch point pay for it once. A CI workflow publishes every main commit's counts as an artifact (``--emit-counts-dir`` is its entry point), and on a disk-cache miss the gate first tries to download the merge-base's artifact through the ``gh`` CLI; any fetch failure falls back silently to the local base pass, so the gate diff --git a/tests/test_litellm/test_default_branch.py b/tests/test_litellm/test_default_branch.py index a1b2a8c5c91..0401bec324f 100644 --- a/tests/test_litellm/test_default_branch.py +++ b/tests/test_litellm/test_default_branch.py @@ -24,7 +24,7 @@ def _commit(repo: Path, message: str) -> None: def remote_and_clone(tmp_path: Path) -> tuple[Path, Path]: seed: Final = tmp_path / "seed" seed.mkdir() - _git(seed, "init", "-q", "-b", "litellm_internal_staging") + _git(seed, "init", "-q", "-b", "release_branch") (seed / "scripts").mkdir() for name in ( "default_branch.py", @@ -47,7 +47,7 @@ def remote_and_clone(tmp_path: Path) -> tuple[Path, Path]: _commit(seed, "main base") remote: Final = tmp_path / "remote.git" _git(tmp_path, "clone", "-q", "--bare", str(seed), str(remote)) - _git(remote, "symbolic-ref", "HEAD", "refs/heads/litellm_internal_staging") + _git(remote, "symbolic-ref", "HEAD", "refs/heads/release_branch") repo: Final = tmp_path / "clone" _git(tmp_path, "clone", "-q", "--single-branch", str(remote), str(repo)) return remote, repo @@ -78,13 +78,13 @@ def test_existing_single_branch_clone_follows_remote_switch(remote_and_clone: tu remote, repo = remote_and_clone before: Final = _resolve(repo) assert before.returncode == 0, before.stderr - assert before.stdout.strip() == "origin/litellm_internal_staging" + assert before.stdout.strip() == "origin/release_branch" _git(remote, "symbolic-ref", "HEAD", "refs/heads/main") after: Final = _resolve(repo) assert after.returncode == 0, after.stderr assert after.stdout.strip() == "origin/main" assert _git(repo, "rev-parse", "origin/main") == _git(remote, "rev-parse", "main") - assert _git(repo, "symbolic-ref", "refs/remotes/origin/HEAD").endswith("/litellm_internal_staging") + assert _git(repo, "symbolic-ref", "refs/remotes/origin/HEAD").endswith("/release_branch") @pytest.mark.parametrize("missing_head", [False, True]) @@ -106,7 +106,7 @@ def test_unverifiable_default_never_uses_cached_head( assert "No changed" not in checked.stdout -@pytest.mark.parametrize("base_ref", ["HEAD", "origin/litellm_internal_staging"]) +@pytest.mark.parametrize("base_ref", ["HEAD", "origin/release_branch"]) def test_explicit_base_works_without_remote_access( remote_and_clone: tuple[Path, Path], base_ref: str, @@ -134,7 +134,7 @@ def test_budget_ratchet_compares_against_new_default(remote_and_clone: tuple[Pat assert "limit raised 0 -> 1" in checked.stdout assert "base origin/main" in checked.stdout overridden: Final = subprocess.run( - [*command, "--base", "origin/litellm_internal_staging"], + [*command, "--base", "origin/release_branch"], cwd=repo, capture_output=True, text=True, @@ -170,7 +170,7 @@ def test_migration_freshness_refuses_stale_branch_after_switch(remote_and_clone: after: Final = _freshness(repo) assert after.returncode == 3 assert "1 commit(s) behind origin/main" in after.stderr - overridden: Final = _freshness(repo, "litellm_internal_staging") + overridden: Final = _freshness(repo, "release_branch") assert overridden.returncode == 0, overridden.stderr _git(repo, "merge", "--ff-only", "origin/main") updated: Final = _freshness(repo) @@ -184,9 +184,9 @@ def test_migration_freshness_refuses_unavailable_remote(remote_and_clone: tuple[ result: Final = _freshness(repo) assert result.returncode == 3 assert "Could not discover origin's default branch" in result.stderr - explicit: Final = _freshness(repo, "litellm_internal_staging") + explicit: Final = _freshness(repo, "release_branch") assert explicit.returncode == 3 - assert "git fetch origin litellm_internal_staging" in explicit.stderr + assert "git fetch origin release_branch" in explicit.stderr @pytest.mark.parametrize( diff --git a/tests/test_litellm/test_git_hooks.py b/tests/test_litellm/test_git_hooks.py index c6980d1f44a..2ecf0da6ed6 100644 --- a/tests/test_litellm/test_git_hooks.py +++ b/tests/test_litellm/test_git_hooks.py @@ -246,7 +246,6 @@ def test_pre_push_rejects_non_conventional_branches(branch): "branch", [ "main", - "litellm_internal_staging", "dependabot/github_actions/foo", "gh-readonly-queue/main/abc123", ], diff --git a/tests/test_litellm/test_pre_commit_lint.py b/tests/test_litellm/test_pre_commit_lint.py index 2d7fa897536..56f98d0e05e 100644 --- a/tests/test_litellm/test_pre_commit_lint.py +++ b/tests/test_litellm/test_pre_commit_lint.py @@ -161,7 +161,7 @@ def _commit_all(repo: Path, message: str) -> None: ) -def _set_base_ref(repo: Path, branch: str = "litellm_internal_staging") -> None: +def _set_base_ref(repo: Path, branch: str = "release_branch") -> None: remote = repo.parent / "remote.git" subprocess.run(["git", "clone", "-q", "--bare", str(repo), str(remote)], check=True) subprocess.run(["git", "update-ref", f"refs/heads/{branch}", "HEAD"], cwd=remote, check=True) @@ -176,7 +176,7 @@ def _stage_file(repo: Path, relative: str, body: str) -> None: subprocess.run(["git", "add", relative], cwd=repo, check=True) -@pytest.mark.parametrize("branch", ["litellm_internal_staging", "main"]) +@pytest.mark.parametrize("branch", ["release_branch", "main"]) def test_nothing_staged_scopes_to_working_tree_diff_and_runs_checks(tmp_path: Path, branch: str) -> None: repo, bin_dir = _sandbox(tmp_path) _commit_all(repo, "base") diff --git a/tests/test_litellm/test_select_ui_test_scope.py b/tests/test_litellm/test_select_ui_test_scope.py index bc11fb495aa..ebfb7701e6a 100644 --- a/tests/test_litellm/test_select_ui_test_scope.py +++ b/tests/test_litellm/test_select_ui_test_scope.py @@ -108,7 +108,7 @@ def _run_step(tmp_path: Path, changed: list[str], base_sha: str = "basesha") -> env["BASE_SHA"] = base_sha env["HEAD_SHA"] = "headsha" env["GITHUB_WORKSPACE"] = str(REPO_ROOT) - env["GITHUB_REF_NAME"] = "litellm_internal_staging" + env["GITHUB_REF_NAME"] = "release_branch" env["CHANGED_FILES"] = str(changed_file) env["NPM_LOG"] = str(npm_log) From 6b642f364878d700f7c9eb98ad69a92868cc4907 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 15:17:36 +0000 Subject: [PATCH 04/10] feat(cost-map): add Azure Foundry pricing for gpt-6-sol and gpt-6-luna (#42747) * feat(cost-map): add Azure Foundry pricing for gpt-6-sol and gpt-6-luna Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(cost-map): give azure/eu gpt-6-sol and gpt-6-luna full model metadata Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 482 ++++++++++++++++++ model_prices_and_context_window.json | 482 ++++++++++++++++++ .../llm_cost_calc/test_llm_cost_calc_utils.py | 28 + .../test_reasoning_effort_capability.py | 27 + 4 files changed, 1019 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 98bda7c12c8..345daa5c8a3 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -3810,6 +3810,104 @@ "supports_web_search": true, "supports_xhigh_reasoning_effort": true }, + "azure_ai/gpt-6-luna": { + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_272k_tokens": 2.5e-07, + "cache_read_input_token_cost": 1e-08, + "cache_read_input_token_cost_above_272k_tokens": 2e-08, + "input_cost_per_token": 1e-07, + "input_cost_per_token_above_272k_tokens": 2e-07, + "litellm_provider": "azure_ai", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "output_cost_per_token_above_272k_tokens": 7.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_cache_breakpoint": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "azure_ai/gpt-6-sol": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "azure_ai", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_cache_breakpoint": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, "azure_ai/gpt-5.5": { "deprecation_date": "2027-10-26", "cache_read_input_token_cost": 5e-07, @@ -7800,6 +7898,198 @@ "supports_web_search": true, "supports_xhigh_reasoning_effort": true }, + "azure/gpt-6-luna": { + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_272k_tokens": 2.5e-07, + "cache_read_input_token_cost": 1e-08, + "cache_read_input_token_cost_above_272k_tokens": 2e-08, + "input_cost_per_token": 1e-07, + "input_cost_per_token_above_272k_tokens": 2e-07, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "output_cost_per_token_above_272k_tokens": 7.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "azure/gpt-6-luna-2026-09-22": { + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_272k_tokens": 2.5e-07, + "cache_read_input_token_cost": 1e-08, + "cache_read_input_token_cost_above_272k_tokens": 2e-08, + "input_cost_per_token": 1e-07, + "input_cost_per_token_above_272k_tokens": 2e-07, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "output_cost_per_token_above_272k_tokens": 7.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "azure/gpt-6-sol": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "azure/gpt-6-sol-2026-09-22": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, "azure/gpt-chat-latest": { "cache_read_input_token_cost": 5e-07, "deprecation_date": "2026-12-02", @@ -8144,6 +8434,102 @@ "supports_web_search": true, "supports_xhigh_reasoning_effort": true }, + "azure/us/gpt-6-luna": { + "cache_creation_input_token_cost": 1.375e-07, + "cache_creation_input_token_cost_above_272k_tokens": 2.75e-07, + "cache_read_input_token_cost": 1.1e-08, + "cache_read_input_token_cost_above_272k_tokens": 2.2e-08, + "input_cost_per_token": 1.1e-07, + "input_cost_per_token_above_272k_tokens": 2.2e-07, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5.5e-07, + "output_cost_per_token_above_272k_tokens": 8.25e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "azure/us/gpt-6-sol": { + "cache_creation_input_token_cost": 2.75e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5.5e-06, + "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost_above_272k_tokens": 4.4e-07, + "input_cost_per_token": 2.2e-06, + "input_cost_per_token_above_272k_tokens": 4.4e-06, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1.1e-05, + "output_cost_per_token_above_272k_tokens": 1.65e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, "azure/us/gpt-chat-latest": { "cache_read_input_token_cost": 5.5e-07, "deprecation_date": "2026-12-02", @@ -66754,6 +67140,102 @@ "output_cost_per_token_above_272k_tokens": 8.25e-05, "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" }, + "azure/eu/gpt-6-luna": { + "cache_creation_input_token_cost": 1.5e-07, + "cache_creation_input_token_cost_above_272k_tokens": 3e-07, + "cache_read_input_token_cost": 1.2e-08, + "cache_read_input_token_cost_above_272k_tokens": 2.4e-08, + "input_cost_per_token": 1.2e-07, + "input_cost_per_token_above_272k_tokens": 2.4e-07, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 6e-07, + "output_cost_per_token_above_272k_tokens": 9e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "azure/eu/gpt-6-sol": { + "cache_creation_input_token_cost": 3e-06, + "cache_creation_input_token_cost_above_272k_tokens": 6e-06, + "cache_read_input_token_cost": 2.4e-07, + "cache_read_input_token_cost_above_272k_tokens": 4.8e-07, + "input_cost_per_token": 2.4e-06, + "input_cost_per_token_above_272k_tokens": 4.8e-06, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1.2e-05, + "output_cost_per_token_above_272k_tokens": 1.8e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, "azure/eu/o1-mini": { "cache_read_input_token_cost": 6.05e-07, "input_cost_per_token": 1.21e-06, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 98bda7c12c8..345daa5c8a3 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -3810,6 +3810,104 @@ "supports_web_search": true, "supports_xhigh_reasoning_effort": true }, + "azure_ai/gpt-6-luna": { + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_272k_tokens": 2.5e-07, + "cache_read_input_token_cost": 1e-08, + "cache_read_input_token_cost_above_272k_tokens": 2e-08, + "input_cost_per_token": 1e-07, + "input_cost_per_token_above_272k_tokens": 2e-07, + "litellm_provider": "azure_ai", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "output_cost_per_token_above_272k_tokens": 7.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_cache_breakpoint": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "azure_ai/gpt-6-sol": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "azure_ai", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_cache_breakpoint": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, "azure_ai/gpt-5.5": { "deprecation_date": "2027-10-26", "cache_read_input_token_cost": 5e-07, @@ -7800,6 +7898,198 @@ "supports_web_search": true, "supports_xhigh_reasoning_effort": true }, + "azure/gpt-6-luna": { + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_272k_tokens": 2.5e-07, + "cache_read_input_token_cost": 1e-08, + "cache_read_input_token_cost_above_272k_tokens": 2e-08, + "input_cost_per_token": 1e-07, + "input_cost_per_token_above_272k_tokens": 2e-07, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "output_cost_per_token_above_272k_tokens": 7.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "azure/gpt-6-luna-2026-09-22": { + "cache_creation_input_token_cost": 1.25e-07, + "cache_creation_input_token_cost_above_272k_tokens": 2.5e-07, + "cache_read_input_token_cost": 1e-08, + "cache_read_input_token_cost_above_272k_tokens": 2e-08, + "input_cost_per_token": 1e-07, + "input_cost_per_token_above_272k_tokens": 2e-07, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5e-07, + "output_cost_per_token_above_272k_tokens": 7.5e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "azure/gpt-6-sol": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "azure/gpt-6-sol-2026-09-22": { + "cache_creation_input_token_cost": 2.5e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5e-06, + "cache_read_input_token_cost": 2e-07, + "cache_read_input_token_cost_above_272k_tokens": 4e-07, + "input_cost_per_token": 2e-06, + "input_cost_per_token_above_272k_tokens": 4e-06, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1e-05, + "output_cost_per_token_above_272k_tokens": 1.5e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, "azure/gpt-chat-latest": { "cache_read_input_token_cost": 5e-07, "deprecation_date": "2026-12-02", @@ -8144,6 +8434,102 @@ "supports_web_search": true, "supports_xhigh_reasoning_effort": true }, + "azure/us/gpt-6-luna": { + "cache_creation_input_token_cost": 1.375e-07, + "cache_creation_input_token_cost_above_272k_tokens": 2.75e-07, + "cache_read_input_token_cost": 1.1e-08, + "cache_read_input_token_cost_above_272k_tokens": 2.2e-08, + "input_cost_per_token": 1.1e-07, + "input_cost_per_token_above_272k_tokens": 2.2e-07, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 5.5e-07, + "output_cost_per_token_above_272k_tokens": 8.25e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "azure/us/gpt-6-sol": { + "cache_creation_input_token_cost": 2.75e-06, + "cache_creation_input_token_cost_above_272k_tokens": 5.5e-06, + "cache_read_input_token_cost": 2.2e-07, + "cache_read_input_token_cost_above_272k_tokens": 4.4e-07, + "input_cost_per_token": 2.2e-06, + "input_cost_per_token_above_272k_tokens": 4.4e-06, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1.1e-05, + "output_cost_per_token_above_272k_tokens": 1.65e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, "azure/us/gpt-chat-latest": { "cache_read_input_token_cost": 5.5e-07, "deprecation_date": "2026-12-02", @@ -66754,6 +67140,102 @@ "output_cost_per_token_above_272k_tokens": 8.25e-05, "source": "https://prices.azure.com/api/retail/prices?$filter=serviceName%20eq%20'Foundry%20Models'%20and%20armRegionName%20eq%20'eastus'%20and%20priceType%20eq%20'Consumption'" }, + "azure/eu/gpt-6-luna": { + "cache_creation_input_token_cost": 1.5e-07, + "cache_creation_input_token_cost_above_272k_tokens": 3e-07, + "cache_read_input_token_cost": 1.2e-08, + "cache_read_input_token_cost_above_272k_tokens": 2.4e-08, + "input_cost_per_token": 1.2e-07, + "input_cost_per_token_above_272k_tokens": 2.4e-07, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 6e-07, + "output_cost_per_token_above_272k_tokens": 9e-07, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, + "azure/eu/gpt-6-sol": { + "cache_creation_input_token_cost": 3e-06, + "cache_creation_input_token_cost_above_272k_tokens": 6e-06, + "cache_read_input_token_cost": 2.4e-07, + "cache_read_input_token_cost_above_272k_tokens": 4.8e-07, + "input_cost_per_token": 2.4e-06, + "input_cost_per_token_above_272k_tokens": 4.8e-06, + "litellm_provider": "azure", + "max_input_tokens": 922000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 1.2e-05, + "output_cost_per_token_above_272k_tokens": 1.8e-05, + "search_context_cost_per_query": { + "search_context_size_high": 0.01, + "search_context_size_low": 0.01, + "search_context_size_medium": 0.01 + }, + "source": "https://azure.microsoft.com/en-us/blog/gpt-6-astra-sol-and-luna-for-production-agents-in-microsoft-foundry/", + "supported_endpoints": [ + "/v1/chat/completions", + "/v1/responses" + ], + "supported_modalities": [ + "text", + "image" + ], + "supported_output_modalities": [ + "text" + ], + "supports_computer_use": true, + "supports_function_calling": true, + "supports_max_reasoning_effort": true, + "supports_minimal_reasoning_effort": false, + "supports_native_streaming": true, + "supports_none_reasoning_effort": true, + "supports_parallel_function_calling": true, + "supports_pdf_input": true, + "supports_prompt_caching": true, + "supports_reasoning": true, + "supports_response_schema": true, + "supports_system_messages": true, + "supports_tool_choice": true, + "supports_vision": true, + "supports_web_search": true, + "supports_xhigh_reasoning_effort": true + }, "azure/eu/o1-mini": { "cache_read_input_token_cost": 6.05e-07, "input_cost_per_token": 1.21e-06, diff --git a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py index 76ccdec25d0..781a3a7c4ed 100644 --- a/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py +++ b/tests/test_litellm/litellm_core_utils/llm_cost_calc/test_llm_cost_calc_utils.py @@ -3761,3 +3761,31 @@ def test_get_batch_cost_rates_has_no_cache_write_rate_without_a_cache_write_batc ) assert rates.cache_creation is None + + +@pytest.mark.parametrize("model_base", ["gpt-6-sol", "gpt-6-luna"]) +def test_azure_gpt_6_foundry_price_sheet(_local_model_cost_map, model_base): + """Azure Foundry hosts gpt-6-sol and gpt-6-luna at OpenAI's Global rates, with the + US and EU data zones charging fixed uplifts on top of them.""" + price_fields = ( + "input_cost_per_token", + "cache_read_input_token_cost", + "cache_creation_input_token_cost", + "output_cost_per_token", + "input_cost_per_token_above_272k_tokens", + "cache_read_input_token_cost_above_272k_tokens", + "cache_creation_input_token_cost_above_272k_tokens", + "output_cost_per_token_above_272k_tokens", + ) + openai_info = litellm.get_model_info(model=model_base, custom_llm_provider="openai") + azure_info = litellm.get_model_info(model=f"azure/{model_base}", custom_llm_provider="azure") + azure_us_info = litellm.get_model_info(model=f"azure/us/{model_base}", custom_llm_provider="azure") + azure_eu_info = litellm.get_model_info(model=f"azure/eu/{model_base}", custom_llm_provider="azure") + azure_ai_info = litellm.get_model_info(model=f"azure_ai/{model_base}", custom_llm_provider="azure_ai") + + for field in price_fields: + base = openai_info[field] + assert azure_info[field] == base + assert azure_ai_info[field] == base + assert azure_us_info[field] == pytest.approx(1.1 * base) + assert azure_eu_info[field] == pytest.approx(1.2 * base) diff --git a/tests/test_litellm/router_utils/test_reasoning_effort_capability.py b/tests/test_litellm/router_utils/test_reasoning_effort_capability.py index 4617839c5e3..2f23320a8ad 100644 --- a/tests/test_litellm/router_utils/test_reasoning_effort_capability.py +++ b/tests/test_litellm/router_utils/test_reasoning_effort_capability.py @@ -458,3 +458,30 @@ class TestNearestDeclaredReasoningEffort: def test_a_level_outside_the_strength_order_is_left_for_upstream(self): assert nearest_declared_reasoning_effort("turbo", ("none", "high")) == "turbo" assert nearest_declared_reasoning_effort("medium", ()) == "medium" + + +class TestAzureGpt6SolAndLunaAdvertiseTheOpenAiLevels: + @pytest.mark.parametrize( + "model,custom_llm_provider", + [ + ("azure/gpt-6-sol", "azure"), + ("azure/gpt-6-luna", "azure"), + ("azure/eu/gpt-6-sol", "azure"), + ("azure/eu/gpt-6-luna", "azure"), + ("azure_ai/gpt-6-sol", "azure_ai"), + ("azure_ai/gpt-6-luna", "azure_ai"), + ], + ) + def test_the_azure_entry_advertises_the_same_levels_as_openai( + self, local_model_cost_map, model, custom_llm_provider + ): + """The Foundry deployments of sol and luna take the same effort set OpenAI documents for + the direct API, so the resolved levels must match the OpenAI-direct entry.""" + from litellm.utils import _get_model_info_helper + + azure_info = dict(_get_model_info_helper(model=model, custom_llm_provider=custom_llm_provider)) + openai_info = dict(_get_model_info_helper(model=model.rsplit("/", 1)[1], custom_llm_provider="openai")) + + assert resolve_supported_reasoning_efforts( + azure_info, deployment_is_mapped=True + ) == resolve_supported_reasoning_efforts(openai_info, deployment_is_mapped=True) From 19556952d93c209be48cac50d8f1b89669a164e9 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 08:24:57 -0700 Subject: [PATCH 05/10] feat(secrets): route secret resolution through native Rust backends (#42619) * fix(secrets): verify provider API request and payload contracts * wip * fix(secrets): unify backend reads and route secret resolution * feat(secrets): bind built-in managers to retained Rust backends * refactor(secrets): centralize catalog dispatch and native binding * test(secrets): split provider integration tests * refactor(secrets): enforce cache and rotation contracts * test(secrets): stub parent packages in failing resolver fixture Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * refactor(secrets): pass manager settings through the interop boundary Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(secrets): align cloud KMS auth and harden provider reads * ci(rust): raise native wheel size gate to 40 MB for secrets backends Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(proxy): treat unset google kms flag as disabled like the old loader Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(secrets): preserve certificate credentials and disabled KMS flags * test(secrets): cover certificate validation and bounded auth retries * test(secrets): cover Python dispatch without the native extension * test(proxy): skip legacy secret manager cases when the optional SDK is missing Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(secrets): port Python parity tests and preserve provider behavior * fix(secrets): store the captured native config without setattr to satisfy the strict lint budget Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(secrets): preserve missing Azure manager values Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(secrets): pin typed values and recovery failure precedence * refactor(secrets): organize provider internals and behavioral test suites * refactor(secrets): simplify recovery and isolate Python compatibility * fix(secrets): distinguish Azure callback absence from HTTP not found * fix(secrets): preserve Python AWS read results at the bridge * fix(secrets): route public reads through the native catalog bridge * fix(secrets): keep JSON selection outside the bridge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(secrets): preserve provider JSON reads at the bridge Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(secrets): preserve Python primary JSON semantics Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(secrets): preserve CyberArk mutation behavior through the native bridge * docs(secrets): record public API replacement gaps * refactor(secrets): share Vault write payload preparation * feat(secrets): route Vault mutations through the native bridge * fix(secrets): preserve typed Vault rotation failures * refactor(secrets): move Python dispatch into bridge * refactor(secrets): move CyberArk Python policy into bridge * refactor(secrets): move Vault Python policy into bridge * test(secrets): assert Vault rotation request paths * fix(secrets): keep bridge JSON interop centralized Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Yujong Lee Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .github/scripts/verify_linux_native_wheel.py | 2 +- litellm-rust/Cargo.lock | 193 +- litellm-rust/crates/core/Cargo.toml | 2 +- litellm-rust/crates/core/src/ocr/prepare.rs | 10 +- litellm-rust/crates/core/tests/ocr.rs | 26 +- litellm-rust/crates/host-python/src/lib.rs | 4 +- .../crates/host-python/src/marshal.rs | 14 + .../crates/llms/src/base_llm/inference/mod.rs | 1 - .../llms/src/base_llm/inference/secrets.rs | 19 - litellm-rust/crates/llms/src/base_llm/mod.rs | 1 - .../crates/llms/src/base_llm/ocr/handler.rs | 4 +- .../llms/src/base_llm/ocr/transformation.rs | 12 +- litellm-rust/crates/python-bridge/Cargo.toml | 3 +- litellm-rust/crates/python-bridge/README.md | 9 +- litellm-rust/crates/python-bridge/src/lib.rs | 10 +- .../python-bridge/src/routes/ocr/mod.rs | 53 +- .../python-bridge/src/secrets/callback.rs | 334 ++-- .../python-bridge/src/secrets/config.rs | 42 +- .../crates/python-bridge/src/secrets/error.rs | 6 +- .../crates/python-bridge/src/secrets/mod.rs | 99 ++ .../python-bridge/src/secrets/mutation.rs | 113 ++ .../python-bridge/src/secrets/operations.rs | 247 +++ .../python-bridge/src/secrets/provider.rs | 150 ++ .../python-bridge/src/secrets/resolved.rs | 124 +- .../python-bridge/src/secrets/runtime.rs | 463 +++++ .../crates/python-bridge/src/secrets/vault.rs | 182 ++ .../src/secrets/vault/operation.rs | 168 ++ litellm-rust/crates/secrets-aws/AGENTS.md | 1 + litellm-rust/crates/secrets-aws/Cargo.toml | 1 + litellm-rust/crates/secrets-aws/src/auth.rs | 23 +- litellm-rust/crates/secrets-aws/src/error.rs | 8 +- litellm-rust/crates/secrets-aws/src/kms.rs | 9 +- .../crates/secrets-aws/src/secret_manager.rs | 402 +---- .../secrets-aws/src/secret_manager/client.rs | 116 ++ .../secrets-aws/src/secret_manager/read.rs | 225 +++ .../secrets-aws/src/secret_manager/write.rs | 329 ++++ litellm-rust/crates/secrets-aws/tests/kms.rs | 11 +- .../secrets-aws/tests/secret_manager.rs | 468 +---- .../tests/secret_manager/configuration.rs | 256 +++ .../secrets-aws/tests/secret_manager/reads.rs | 198 +++ .../tests/secret_manager/support.rs | 80 + .../tests/secret_manager/writes.rs | 615 +++++++ litellm-rust/crates/secrets-azure/AGENTS.md | 1 + litellm-rust/crates/secrets-azure/Cargo.toml | 2 +- .../crates/secrets-azure/src/error.rs | 4 + .../crates/secrets-azure/src/key_vault.rs | 67 +- litellm-rust/crates/secrets-azure/src/lib.rs | 2 +- .../crates/secrets-azure/tests/key_vault.rs | 37 + .../crates/secrets-cyberark/AGENTS.md | 1 + .../crates/secrets-cyberark/Cargo.toml | 2 + .../crates/secrets-cyberark/src/error.rs | 2 + .../crates/secrets-cyberark/src/lib.rs | 2 +- .../secrets-cyberark/src/secret_manager.rs | 359 +--- .../src/secret_manager/client.rs | 147 ++ .../src/secret_manager/read.rs | 118 ++ .../src/secret_manager/write.rs | 262 +++ .../secrets-cyberark/tests/secret_manager.rs | 563 +----- .../tests/secret_manager/configuration.rs | 452 +++++ .../tests/secret_manager/reads.rs | 145 ++ .../tests/secret_manager/support.rs | 70 + .../tests/secret_manager/writes.rs | 450 +++++ litellm-rust/crates/secrets-google/AGENTS.md | 1 + litellm-rust/crates/secrets-google/Cargo.toml | 5 +- .../crates/secrets-google/src/error.rs | 6 + litellm-rust/crates/secrets-google/src/kms.rs | 21 +- .../secrets-google/src/secret_manager.rs | 86 +- .../crates/secrets-google/tests/kms.rs | 43 +- .../secrets-google/tests/secret_manager.rs | 163 ++ .../crates/secrets-hashicorp/AGENTS.md | 1 + .../crates/secrets-hashicorp/Cargo.toml | 3 +- .../crates/secrets-hashicorp/src/error.rs | 10 +- .../crates/secrets-hashicorp/src/lib.rs | 2 +- .../secrets-hashicorp/src/secret_manager.rs | 341 +--- .../src/secret_manager/client.rs | 115 ++ .../src/secret_manager/raw.rs | 155 ++ .../src/secret_manager/read.rs | 63 + .../src/secret_manager/write.rs | 190 ++ .../secrets-hashicorp/tests/secret_manager.rs | 932 +--------- .../tests/secret_manager/configuration.rs | 335 ++++ .../tests/secret_manager/reads.rs | 385 ++++ .../tests/secret_manager/support.rs | 175 ++ .../tests/secret_manager/writes.rs | 542 ++++++ litellm-rust/crates/secrets-types/Cargo.toml | 2 + .../secrets-types/src/base_secret_manager.rs | 130 +- .../crates/secrets-types/src/cache.rs | 73 + .../crates/secrets-types/src/context.rs | 45 +- .../crates/secrets-types/src/error.rs | 4 + litellm-rust/crates/secrets-types/src/lib.rs | 13 +- .../crates/secrets-types/src/value.rs | 6 + .../crates/secrets-types/tests/cache.rs | 162 ++ .../crates/secrets-types/tests/context.rs | 42 +- .../crates/secrets-types/tests/rotation.rs | 165 +- litellm-rust/crates/secrets/Cargo.toml | 3 +- litellm-rust/crates/secrets/PARITY.md | 277 +++ litellm-rust/crates/secrets/README.md | 38 +- .../crates/secrets/src/compatibility.rs | 82 + litellm-rust/crates/secrets/src/error.rs | 6 + litellm-rust/crates/secrets/src/handler.rs | 49 +- litellm-rust/crates/secrets/src/lib.rs | 7 +- litellm-rust/crates/secrets/src/native.rs | 61 + litellm-rust/crates/secrets/src/oidc.rs | 41 +- litellm-rust/crates/secrets/src/resolver.rs | 130 +- litellm-rust/crates/secrets/src/source.rs | 73 + litellm-rust/crates/secrets/tests/aws.rs | 223 +++ litellm-rust/crates/secrets/tests/azure.rs | 106 ++ .../secrets/tests/common_read_contract.rs | 260 +++ litellm-rust/crates/secrets/tests/cyberark.rs | 53 + litellm-rust/crates/secrets/tests/google.rs | 117 ++ litellm-rust/crates/secrets/tests/handler.rs | 363 ---- .../crates/secrets/tests/hashicorp.rs | 143 ++ litellm-rust/crates/secrets/tests/oidc.rs | 195 +- .../crates/secrets/tests/resolution.rs | 553 +++--- litellm-rust/crates/secrets/tests/source.rs | 62 + litellm/proxy/proxy_server.py | 12 +- litellm/rust_bridge/_native.pyi | 42 + litellm/rust_bridge/secret_manager.py | 285 +++ litellm/rust_bridge/settings.py | 22 +- .../secret_managers/aws_secret_manager_v2.py | 9 + .../cyberark_secret_manager.py | 44 + litellm/secret_managers/dispatch.py | 30 + .../secret_managers/google_secret_manager.py | 5 + .../hashicorp_secret_manager.py | 25 + litellm/secret_managers/main.py | 2 +- .../proxy/proxy_server/test_proxy_config.py | 83 + tests/test_litellm/rust_bridge/AGENTS.md | 7 + .../rust_bridge/ocr/test_secrets.py | 459 ++++- .../rust_bridge/test_secret_manager.py | 1571 +++++++++++++++++ .../test_litellm/rust_bridge/test_settings.py | 45 +- .../support/recording_server.py | 3 + 129 files changed, 13275 insertions(+), 4146 deletions(-) delete mode 100644 litellm-rust/crates/llms/src/base_llm/inference/mod.rs delete mode 100644 litellm-rust/crates/llms/src/base_llm/inference/secrets.rs create mode 100644 litellm-rust/crates/python-bridge/src/secrets/mutation.rs create mode 100644 litellm-rust/crates/python-bridge/src/secrets/operations.rs create mode 100644 litellm-rust/crates/python-bridge/src/secrets/provider.rs create mode 100644 litellm-rust/crates/python-bridge/src/secrets/runtime.rs create mode 100644 litellm-rust/crates/python-bridge/src/secrets/vault.rs create mode 100644 litellm-rust/crates/python-bridge/src/secrets/vault/operation.rs create mode 100644 litellm-rust/crates/secrets-aws/AGENTS.md create mode 100644 litellm-rust/crates/secrets-aws/src/secret_manager/client.rs create mode 100644 litellm-rust/crates/secrets-aws/src/secret_manager/read.rs create mode 100644 litellm-rust/crates/secrets-aws/src/secret_manager/write.rs create mode 100644 litellm-rust/crates/secrets-aws/tests/secret_manager/configuration.rs create mode 100644 litellm-rust/crates/secrets-aws/tests/secret_manager/reads.rs create mode 100644 litellm-rust/crates/secrets-aws/tests/secret_manager/support.rs create mode 100644 litellm-rust/crates/secrets-aws/tests/secret_manager/writes.rs create mode 100644 litellm-rust/crates/secrets-azure/AGENTS.md create mode 100644 litellm-rust/crates/secrets-cyberark/AGENTS.md create mode 100644 litellm-rust/crates/secrets-cyberark/src/secret_manager/client.rs create mode 100644 litellm-rust/crates/secrets-cyberark/src/secret_manager/read.rs create mode 100644 litellm-rust/crates/secrets-cyberark/src/secret_manager/write.rs create mode 100644 litellm-rust/crates/secrets-cyberark/tests/secret_manager/configuration.rs create mode 100644 litellm-rust/crates/secrets-cyberark/tests/secret_manager/reads.rs create mode 100644 litellm-rust/crates/secrets-cyberark/tests/secret_manager/support.rs create mode 100644 litellm-rust/crates/secrets-cyberark/tests/secret_manager/writes.rs create mode 100644 litellm-rust/crates/secrets-google/AGENTS.md create mode 100644 litellm-rust/crates/secrets-hashicorp/AGENTS.md create mode 100644 litellm-rust/crates/secrets-hashicorp/src/secret_manager/client.rs create mode 100644 litellm-rust/crates/secrets-hashicorp/src/secret_manager/raw.rs create mode 100644 litellm-rust/crates/secrets-hashicorp/src/secret_manager/read.rs create mode 100644 litellm-rust/crates/secrets-hashicorp/src/secret_manager/write.rs create mode 100644 litellm-rust/crates/secrets-hashicorp/tests/secret_manager/configuration.rs create mode 100644 litellm-rust/crates/secrets-hashicorp/tests/secret_manager/reads.rs create mode 100644 litellm-rust/crates/secrets-hashicorp/tests/secret_manager/support.rs create mode 100644 litellm-rust/crates/secrets-hashicorp/tests/secret_manager/writes.rs create mode 100644 litellm-rust/crates/secrets-types/src/cache.rs create mode 100644 litellm-rust/crates/secrets-types/tests/cache.rs create mode 100644 litellm-rust/crates/secrets/PARITY.md create mode 100644 litellm-rust/crates/secrets/src/compatibility.rs create mode 100644 litellm-rust/crates/secrets/src/native.rs create mode 100644 litellm-rust/crates/secrets/src/source.rs create mode 100644 litellm-rust/crates/secrets/tests/aws.rs create mode 100644 litellm-rust/crates/secrets/tests/azure.rs create mode 100644 litellm-rust/crates/secrets/tests/common_read_contract.rs create mode 100644 litellm-rust/crates/secrets/tests/cyberark.rs create mode 100644 litellm-rust/crates/secrets/tests/google.rs delete mode 100644 litellm-rust/crates/secrets/tests/handler.rs create mode 100644 litellm-rust/crates/secrets/tests/hashicorp.rs create mode 100644 litellm-rust/crates/secrets/tests/source.rs create mode 100644 litellm/rust_bridge/secret_manager.py create mode 100644 litellm/secret_managers/dispatch.py create mode 100644 tests/test_litellm/rust_bridge/AGENTS.md create mode 100644 tests/test_litellm/rust_bridge/test_secret_manager.py diff --git a/.github/scripts/verify_linux_native_wheel.py b/.github/scripts/verify_linux_native_wheel.py index f2b82f86b47..6b7fcd57bbc 100644 --- a/.github/scripts/verify_linux_native_wheel.py +++ b/.github/scripts/verify_linux_native_wheel.py @@ -214,7 +214,7 @@ def main( native_module: Final = load_native_module(native_path) native_module_loads: Final = native_module is not None panic_test_hook_absent: Final = native_module is not None and not hasattr(native_module, "_panic_for_test") - native_size_limit: Final = 35_000_000 + native_size_limit: Final = 40_000_000 native_size_within_limit: Final = native_member.file_size <= native_size_limit validations: Final = ( (f"Python tag is {EXPECTED_PYTHON_TAG}", python_tag == EXPECTED_PYTHON_TAG), diff --git a/litellm-rust/Cargo.lock b/litellm-rust/Cargo.lock index 332d8b3dcf0..ef032bfe55c 100644 --- a/litellm-rust/Cargo.lock +++ b/litellm-rust/Cargo.lock @@ -88,6 +88,45 @@ version = "1.2.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "03918c3dbd7701a85c6b9887732e2921175f26c350b4563841d0958c21d57e6d" +[[package]] +name = "asn1-rs" +version = "0.7.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b7f43a50ac4fdca5df8e885c21b835997f0a1cdee65494a6847694a98652d9d8" +dependencies = [ + "asn1-rs-derive", + "asn1-rs-impl", + "displaydoc", + "nom", + "num-traits", + "rusticata-macros", + "thiserror 2.0.19", + "time", +] + +[[package]] +name = "asn1-rs-derive" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3109e49b1e4909e9db6515a30c633684d68cdeaa252f215214cb4fa1a5bfee2c" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", + "synstructure 0.13.2", +] + +[[package]] +name = "asn1-rs-impl" +version = "0.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7b18050c2cd6fe86c3a76584ef5e0baf286d038cda203eb6223df2cc413565f7" +dependencies = [ + "proc-macro2", + "quote", + "syn 2.0.119", +] + [[package]] name = "assert-json-diff" version = "2.0.2" @@ -810,7 +849,7 @@ version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "08807e080ed7f9d5433fa9b275196cfc35414f66a0c79d864dc51a0d825231a3" dependencies = [ - "bit-vec", + "bit-vec 0.8.0", ] [[package]] @@ -819,6 +858,15 @@ version = "0.8.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5e764a1d40d510daf35e07be9eb06e75770908c27d411ee6c92109c9840eaaf7" +[[package]] +name = "bit-vec" +version = "0.9.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b71798fca2c1fe1086445a7258a4bc81e6e49dcd24c8d0dd9a1e57395b603f51" +dependencies = [ + "serde", +] + [[package]] name = "bitflags" version = "1.3.2" @@ -1120,6 +1168,15 @@ version = "0.4.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "338089f42c427b86394a5ee60ff321da23a5c89c9d89514c829687b26359fcff" +[[package]] +name = "crc32c" +version = "0.6.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3a47af21622d091a8f0fb295b88bc886ac74efcc613efc19f5d0b21de5c89e47" +dependencies = [ + "rustc_version", +] + [[package]] name = "crc32fast" version = "1.5.1" @@ -1378,6 +1435,20 @@ dependencies = [ "thiserror 2.0.19", ] +[[package]] +name = "der-parser" +version = "10.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "07da5016415d5a3c4dd39b11ed26f915f52fc4e0dc197d87908bc916e51bc1a6" +dependencies = [ + "asn1-rs", + "displaydoc", + "nom", + "num-bigint 0.4.8", + "num-traits", + "rusticata-macros", +] + [[package]] name = "deranged" version = "0.5.8" @@ -2589,21 +2660,6 @@ dependencies = [ "wasm-bindgen", ] -[[package]] -name = "jsonwebtoken" -version = "11.1.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "e75fe14a82d81e5f5af639997db37d8b96045938a7ac6ab18cdbe1c7467e05e1" -dependencies = [ - "base64 0.22.1", - "getrandom 0.2.17", - "js-sys", - "serde", - "serde_json", - "signature", - "zeroize", -] - [[package]] name = "lazy_static" version = "1.5.0" @@ -3119,6 +3175,7 @@ dependencies = [ "tokio", "tokio-tungstenite", "url", + "veil", "wiremock", ] @@ -3143,10 +3200,11 @@ version = "0.1.0" dependencies = [ "aws-sdk-kms", "base64 0.22.1", + "futures-util", "google-cloud-auth", "google-cloud-kms-v1", - "jsonwebtoken", "litellm-core-utils", + "litellm-python-compat", "litellm-secrets-aws", "litellm-secrets-azure", "litellm-secrets-cyberark", @@ -3179,6 +3237,7 @@ dependencies = [ "litellm-tracing", "rstest", "serde_json", + "tempfile", "thiserror 2.0.19", "tokio", "veil", @@ -3215,10 +3274,12 @@ dependencies = [ "litellm-tracing", "moka", "percent-encoding", + "rcgen", "reqwest 0.12.28", "rstest", "serde", "serde_json", + "tempfile", "thiserror 2.0.19", "tokio", "veil", @@ -3230,6 +3291,7 @@ name = "litellm-secrets-google" version = "0.1.0" dependencies = [ "base64 0.22.1", + "crc32c", "google-cloud-auth", "google-cloud-gax", "google-cloud-kms-v1", @@ -3255,7 +3317,6 @@ version = "0.1.0" dependencies = [ "litellm-core-utils", "litellm-secrets-types", - "moka", "rstest", "rustify", "rustify_derive", @@ -3274,6 +3335,7 @@ name = "litellm-secrets-types" version = "0.1.0" dependencies = [ "litellm-auth-types", + "moka", "rstest", "serde", "serde_json", @@ -3588,6 +3650,15 @@ dependencies = [ "libc", ] +[[package]] +name = "oid-registry" +version = "0.8.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "12f40cff3dde1b6087cc5d5f5d4d65712f34016a03ed60e9c08dcc392736b5b7" +dependencies = [ + "asn1-rs", +] + [[package]] name = "once_cell" version = "1.21.4" @@ -3721,6 +3792,16 @@ version = "0.2.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2ee67f1008b1ba2321834326597b8e186293b049a023cdef258527550b9935b4" +[[package]] +name = "pem" +version = "4.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d354a98a3d1251555de99e8fdd8afda05573c31b82f59063a7b0a29b5527f120" +dependencies = [ + "base64 0.23.1", + "serde_core", +] + [[package]] name = "percent-encoding" version = "2.3.2" @@ -3899,7 +3980,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "4b45fcc2344c680f5025fe57779faef368840d0bd1f42f216291f0dc4ace4744" dependencies = [ "bit-set", - "bit-vec", + "bit-vec 0.8.0", "bitflags 2.13.1", "num-traits", "rand 0.9.5", @@ -4288,6 +4369,20 @@ dependencies = [ "crossbeam-utils", ] +[[package]] +name = "rcgen" +version = "0.14.10" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8774e05a7d0de114588e6a28fe7e71694b82614ed569d86d8b389dfbc98b8ad8" +dependencies = [ + "pem", + "ring", + "rustls-pki-types", + "time", + "x509-parser", + "yasna", +] + [[package]] name = "redis" version = "1.7.0" @@ -4570,6 +4665,15 @@ dependencies = [ "semver", ] +[[package]] +name = "rusticata-macros" +version = "4.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "faf0c4a6ece9950b9abdb62b1cfcf2a68b3b67a10ba445b3bb85be2a293d0632" +dependencies = [ + "nom", +] + [[package]] name = "rustify" version = "0.7.0" @@ -5046,15 +5150,6 @@ dependencies = [ "libc", ] -[[package]] -name = "signature" -version = "2.2.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "77549399552de45a898a580c1b41d445bf730df867cc44e6c0233bbc4b8329de" -dependencies = [ - "rand_core 0.6.4", -] - [[package]] name = "simd-adler32" version = "0.3.10" @@ -6356,6 +6451,24 @@ version = "0.6.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1ffae5123b2d3fc086436f8834ae3ab053a283cfac8fe0a0b8eaae044768a4c4" +[[package]] +name = "x509-parser" +version = "0.18.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d43b0f71ce057da06bc0851b23ee24f3f86190b07203dd8f567d0b706a185202" +dependencies = [ + "asn1-rs", + "data-encoding", + "der-parser", + "lazy_static", + "nom", + "oid-registry", + "ring", + "rusticata-macros", + "thiserror 2.0.19", + "time", +] + [[package]] name = "xmlparser" version = "0.13.6" @@ -6368,6 +6481,16 @@ version = "0.8.18" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "aee1b19627c7c60102ab80d3a9cbe18de90bfe03bfa6c3715447681f0e8c8af6" +[[package]] +name = "yasna" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b5f6765e852b9b4dc8e2a76843e4d64d1cea8e79bcde0b6901aea8e7c7f08282" +dependencies = [ + "bit-vec 0.9.1", + "time", +] + [[package]] name = "yoke" version = "0.8.3" @@ -6437,20 +6560,6 @@ name = "zeroize" version = "1.9.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e13c156562582aa81c60cb29407084cdb54c4164760106ab78e6c5b0858cf64e" -dependencies = [ - "zeroize_derive", -] - -[[package]] -name = "zeroize_derive" -version = "1.5.0" -source = "registry+https://github.com/rust-lang/crates.io-index" -checksum = "3c50655cbb0fe3fc43170059e702f1ce5e19b84cec58dc87b037a09935c2f328" -dependencies = [ - "proc-macro2", - "quote", - "syn 2.0.119", -] [[package]] name = "zerotrie" diff --git a/litellm-rust/crates/core/Cargo.toml b/litellm-rust/crates/core/Cargo.toml index 3bfc5bae925..4626781f3e3 100644 --- a/litellm-rust/crates/core/Cargo.toml +++ b/litellm-rust/crates/core/Cargo.toml @@ -7,6 +7,7 @@ repository.workspace = true autotests = false [dependencies] +litellm-secrets.workspace = true litellm-types.workspace = true litellm-core-utils.workspace = true litellm-host.workspace = true @@ -36,7 +37,6 @@ url.workspace = true veil.workspace = true [dev-dependencies] -litellm-secrets.workspace = true litellm-auth-gcp.workspace = true litellm-llms = { workspace = true, features = ["test-support"] } rstest.workspace = true diff --git a/litellm-rust/crates/core/src/ocr/prepare.rs b/litellm-rust/crates/core/src/ocr/prepare.rs index 54960256faa..37c0f18f659 100644 --- a/litellm-rust/crates/core/src/ocr/prepare.rs +++ b/litellm-rust/crates/core/src/ocr/prepare.rs @@ -1,11 +1,9 @@ use litellm_auth::{InputSource, SecretValue, Sourced}; -use litellm_llms::base_llm::{ - inference::secrets::Secrets, - ocr::{ - handler::OcrClient, - transformation::{OcrConnection, OcrCredentialInputs, PreparedOcrRequest}, - }, +use litellm_llms::base_llm::ocr::{ + handler::OcrClient, + transformation::{OcrConnection, OcrCredentialInputs, PreparedOcrRequest}, }; +use litellm_secrets::source::Secrets; use super::provider_config::OcrProvider; use crate::ocr::types::{LiteLLMOcrRequest, ResolvedOcrRequest}; diff --git a/litellm-rust/crates/core/tests/ocr.rs b/litellm-rust/crates/core/tests/ocr.rs index d376f0df784..26cd2153e8d 100644 --- a/litellm-rust/crates/core/tests/ocr.rs +++ b/litellm-rust/crates/core/tests/ocr.rs @@ -11,7 +11,6 @@ use litellm_http::{ HttpClientPool, HttpSettings, Resolution, media::{PublicDnsResolver, UrlPolicy}, }; -use litellm_llms::base_llm::inference::secrets::{SecretSource, Secrets}; use litellm_llms::base_llm::ocr::{ error::Error as OcrError, handler::OcrClient, @@ -20,6 +19,7 @@ use litellm_llms::base_llm::ocr::{ BaseOcrConfig, LiteLLMOcrResponse, OCR_RESPONSE_MAX_BYTES, OcrTransportConfig, }, }; +use litellm_secrets::source::SecretSource; use rstest::rstest; use serde_json::{Value, json}; @@ -32,27 +32,27 @@ use super::{ use crate::ocr::route::{LocalOcrHost, OcrOp, OcrOpResult, ocr_machine}; struct RecordingSecretSource { - names: Arc>>, + names: Arc>>, values: &'static [(&'static str, &'static str)], api_base: String, } impl SecretSource for RecordingSecretSource { - fn resolve<'a>( + fn get_secret_str<'a>( &'a self, - names: &'a [&'static str], - ) -> BoxFuture<'a, Result> { - *self.names.lock().unwrap() = names.to_vec(); - let values = self.values; - let api_base = self.api_base.clone(); + name: &'a str, + ) -> BoxFuture<'a, Result, litellm_secrets::Error>> { + self.names.lock().unwrap().push(name.to_owned()); Box::pin(async move { - Ok(Arc::new(move |name: &str| match name { - "MISTRAL_AZURE_API_BASE" => Some(api_base.clone()), - _ => values + Ok(match name { + "MISTRAL_AZURE_API_BASE" => Some(self.api_base.clone()), + _ => self + .values .iter() .find(|(key, _)| *key == name) .map(|(_, value)| value.to_string()), - }) as Secrets) + } + .map(litellm_secrets::SecretValue::new)) }) } } @@ -289,7 +289,7 @@ async fn ocr_client_uses_the_injected_http_pool_configuration() { UrlPolicy::default(), VertexAuth::default(), OcrSettings::default(), - Arc::new(litellm_llms::base_llm::inference::secrets::EnvironmentSecrets), + Arc::new(litellm_secrets::source::EnvironmentSecrets::default()), ) .unwrap(); crate::ocr::client::perform(&client, wire_request("mistral/model", &base, json!({}))) diff --git a/litellm-rust/crates/host-python/src/lib.rs b/litellm-rust/crates/host-python/src/lib.rs index 4a33975a918..37f8a2c9e4b 100644 --- a/litellm-rust/crates/host-python/src/lib.rs +++ b/litellm-rust/crates/host-python/src/lib.rs @@ -27,7 +27,9 @@ pub use execution::{ pub use fork_gate::RuntimeAlreadyStarted; pub use gil::{release_count, release_gil}; pub use handle::{Execution, ExecutionBody, ExecutionStep}; -pub use marshal::{Pythonized, from_py, from_py_argument, panic_to_pyerr, to_py}; +pub use marshal::{ + Pythonized, from_py, from_py_argument, json_loads, json_object_field, panic_to_pyerr, to_py, +}; /// Starts the interpreter and imports the standard modules the tests share, once, so /// parallel test threads never race a first import of `asyncio`. diff --git a/litellm-rust/crates/host-python/src/marshal.rs b/litellm-rust/crates/host-python/src/marshal.rs index 8f284abf9dd..53f11d8c40a 100644 --- a/litellm-rust/crates/host-python/src/marshal.rs +++ b/litellm-rust/crates/host-python/src/marshal.rs @@ -4,6 +4,7 @@ use std::panic::{AssertUnwindSafe, catch_unwind}; use pyo3::exceptions::PyValueError; use pyo3::panic::PanicException; use pyo3::prelude::*; +use pyo3::types::PyBytes; use serde::Serialize; use serde::de::DeserializeOwned; @@ -32,6 +33,19 @@ where .map_err(PyErr::from) } +pub fn json_object_field(py: Python<'_>, document: &str, name: &str) -> PyResult> { + py.import("json")? + .call_method1("loads", (document,))? + .call_method1("get", (name,)) + .map(Bound::unbind) +} + +pub fn json_loads(py: Python<'_>, document: &[u8]) -> PyResult> { + py.import("json")? + .call_method1("loads", (PyBytes::new(py, document),)) + .map(Bound::unbind) +} + pub struct Pythonized(pub T); impl<'py, T> IntoPyObject<'py> for Pythonized diff --git a/litellm-rust/crates/llms/src/base_llm/inference/mod.rs b/litellm-rust/crates/llms/src/base_llm/inference/mod.rs deleted file mode 100644 index 10c0454f947..00000000000 --- a/litellm-rust/crates/llms/src/base_llm/inference/mod.rs +++ /dev/null @@ -1 +0,0 @@ -pub mod secrets; diff --git a/litellm-rust/crates/llms/src/base_llm/inference/secrets.rs b/litellm-rust/crates/llms/src/base_llm/inference/secrets.rs deleted file mode 100644 index eb13fe95116..00000000000 --- a/litellm-rust/crates/llms/src/base_llm/inference/secrets.rs +++ /dev/null @@ -1,19 +0,0 @@ -use std::sync::Arc; - -use futures_util::future::BoxFuture; -use litellm_core_utils::settings::{Lookup, ProcessEnvironment}; -use litellm_secrets::Error; - -pub type Secrets = Arc; - -pub trait SecretSource: Send + Sync { - fn resolve<'a>(&'a self, names: &'a [&'static str]) -> BoxFuture<'a, Result>; -} - -pub struct EnvironmentSecrets; - -impl SecretSource for EnvironmentSecrets { - fn resolve<'a>(&'a self, _names: &'a [&'static str]) -> BoxFuture<'a, Result> { - Box::pin(async { Ok(Arc::new(ProcessEnvironment) as Secrets) }) - } -} diff --git a/litellm-rust/crates/llms/src/base_llm/mod.rs b/litellm-rust/crates/llms/src/base_llm/mod.rs index 9cced64b687..8ed37da4573 100644 --- a/litellm-rust/crates/llms/src/base_llm/mod.rs +++ b/litellm-rust/crates/llms/src/base_llm/mod.rs @@ -2,6 +2,5 @@ pub mod anthropic_messages; pub mod audio_transcription; pub mod base_model_iterator; pub mod chat; -pub mod inference; pub mod ocr; pub mod responses; diff --git a/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs b/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs index 3ec9de8197f..9527fd20f2d 100644 --- a/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs +++ b/litellm-rust/crates/llms/src/base_llm/ocr/handler.rs @@ -13,7 +13,6 @@ use litellm_http::{ use serde::{Serialize, de::DeserializeOwned}; use serde_json::Value; -use crate::base_llm::inference::secrets::SecretSource; use crate::base_llm::ocr::{ error::Error, settings::OcrSettings, @@ -22,6 +21,7 @@ use crate::base_llm::ocr::{ PreparedOcrRequest, decode_request_value, decode_response, }, }; +use litellm_secrets::source::SecretSource; /// The route's view of one call, handed to provider code that has to reach the /// caller's hooks mid-flight (guardrails on the outgoing body, raw response events). @@ -95,7 +95,7 @@ impl OcrClient { document_fetcher: MediaFetcher::for_test(document_http), vertex_auth: VertexAuth::default(), settings: OcrSettings::default(), - secrets: Arc::new(crate::base_llm::inference::secrets::EnvironmentSecrets), + secrets: Arc::new(litellm_secrets::source::EnvironmentSecrets::default()), } } diff --git a/litellm-rust/crates/llms/src/base_llm/ocr/transformation.rs b/litellm-rust/crates/llms/src/base_llm/ocr/transformation.rs index 0506ff3d6df..3f1b260bb13 100644 --- a/litellm-rust/crates/llms/src/base_llm/ocr/transformation.rs +++ b/litellm-rust/crates/llms/src/base_llm/ocr/transformation.rs @@ -7,6 +7,7 @@ use litellm_core_utils::{ settings::ProcessEnvironment, }; use litellm_http::outbound::{OutboundRequest, RequestSigner}; +use litellm_secrets::source::Secrets; use serde::{ Deserialize, Serialize, de::{DeserializeOwned, IntoDeserializer}, @@ -14,13 +15,10 @@ use serde::{ use serde_json::{Map, Value}; use serde_with::serde_as; -use crate::base_llm::{ - inference::secrets::Secrets, - ocr::{ - error::Error, - handler::{CallHooks, OcrClient, read_response_bytes, transform_request_body}, - settings::OcrSettings, - }, +use crate::base_llm::ocr::{ + error::Error, + handler::{CallHooks, OcrClient, read_response_bytes, transform_request_body}, + settings::OcrSettings, }; pub const OCR_RESPONSE_MAX_BYTES: usize = 64 * 1024 * 1024; diff --git a/litellm-rust/crates/python-bridge/Cargo.toml b/litellm-rust/crates/python-bridge/Cargo.toml index 85e6b9588dd..a02adfaa064 100644 --- a/litellm-rust/crates/python-bridge/Cargo.toml +++ b/litellm-rust/crates/python-bridge/Cargo.toml @@ -45,7 +45,7 @@ litellm-core-utils.workspace = true litellm-auth-gcp.workspace = true litellm-http.workspace = true litellm-llms.workspace = true -litellm-secrets = { workspace = true, features = ["aws"] } +litellm-secrets = { workspace = true, features = ["aws", "azure", "google", "hashicorp", "cyberark"] } litellm-secrets-types.workspace = true litellm-types.workspace = true litellm-host-python.workspace = true @@ -55,6 +55,7 @@ pyo3-async-runtimes.workspace = true reqwest.workspace = true redis = { version = "1.7.0", features = ["tls-rustls"] } serde_json.workspace = true +veil.workspace = true thiserror.workspace = true tokio = { workspace = true, features = ["rt", "sync"] } url.workspace = true diff --git a/litellm-rust/crates/python-bridge/README.md b/litellm-rust/crates/python-bridge/README.md index faaca233f5a..fa8df8757ff 100644 --- a/litellm-rust/crates/python-bridge/README.md +++ b/litellm-rust/crates/python-bridge/README.md @@ -1,5 +1,10 @@ -Native OCR uses `SecretSource` with `EnvironmentSecrets`, preserving process-environment reads. Readable Python secret managers still make OCR decline to the existing Python implementation. `ResolvedSecrets` and the separate `secret_manager_binding()` snapshot are inactive foundations for a later rollout +Native OCR uses `litellm_secrets::source::SecretSource`. Built-in secret managers resolve to retained Rust backends. Custom Python managers and overrides keep the callback path. Readable managers still require the Rust secret-manager binding to be enabled -Cache and secret-manager catalog entries remain Python-only, including when `LITELLM_RUST=1`. The new cache runtime is not connected to SDK or gateway caching +The shared proxy initializer captures native configuration without loading the extension or doing native I/O. `_SecretManagerRuntime.from_client` constructs a backend on first use and keeps its handle on the Python client. The secret-manager dispatcher selects Python or Rust through `catalog.py`. Native reads call that handle; Rust routes extract the backend directly. Configuration changes replace the handle, while calls already bound to the previous backend keep using it. Handles cannot be reused after fork. Directly constructed LiteLLM managers are adapted on first native use. Manually supplied SDK clients keep their Python behavior because their credentials cannot be inferred safely. Provider implementations contain no bridge registration + +Retention describes ownership and lifetime. `callbacks-legacy-python::PublicCall` owns Python references for one call to preserve identity. A native cache or secret-manager handle owns shared Rust state across calls to preserve connection pools and caches. Both use existing `Py` and shared Rust ownership, with execution and GIL transitions handled by `litellm-host-python` + + +Cache and secret-manager catalog entries remain Python-only, including when `LITELLM_RUST=1`. This wiring does not change rollout policy OCR provider requests use the shared `litellm-http` pool. AWS and Google secret-manager SDK clients keep their SDK transports, which do not yet inherit the pool's proxy, TLS, certificate, timeout, or observability configuration. Preserve those SDK transports and configure them equivalently instead of forcing them through reqwest diff --git a/litellm-rust/crates/python-bridge/src/lib.rs b/litellm-rust/crates/python-bridge/src/lib.rs index 76de2ae2b7b..54b13ba01bb 100644 --- a/litellm-rust/crates/python-bridge/src/lib.rs +++ b/litellm-rust/crates/python-bridge/src/lib.rs @@ -8,10 +8,6 @@ mod logger; mod marshal; mod python_settings; mod routes; -#[allow( - dead_code, - reason = "secret-manager foundations await rollout activation" -)] mod secrets; mod tokenizer; @@ -56,7 +52,11 @@ mod _native { let dict = module.dict(); dict.set_item("_CacheTestHandle", py.get_type::())?; dict.set_item("_CacheTestResolver", py.get_type::())?; - dict.set_item("_ResponseCacheRuntime", py.get_type::()) + dict.set_item("_ResponseCacheRuntime", py.get_type::())?; + dict.set_item( + "_SecretManagerRuntime", + py.get_type::(), + ) } } diff --git a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs index ca5abf1f80e..55b6b0ce21d 100644 --- a/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs +++ b/litellm-rust/crates/python-bridge/src/routes/ocr/mod.rs @@ -3,17 +3,14 @@ mod errors; mod host; mod project; -use std::sync::{Arc, LazyLock}; +use std::sync::LazyLock; use host::OcrRouteHost; use litellm_auth_gcp::VertexAuth; use litellm_callbacks_legacy_python::{LegacySurface, PublicCall, run_legacy_call}; use litellm_core::ocr::route::ocr_machine; use litellm_core_utils::settings::ProcessEnvironment; -use litellm_llms::base_llm::{ - inference::secrets::{EnvironmentSecrets, SecretSource}, - ocr::{handler::OcrClient, settings::OcrSettings}, -}; +use litellm_llms::base_llm::ocr::{handler::OcrClient, settings::OcrSettings}; use pyo3::{ prelude::*, types::{PyDict, PyTuple}, @@ -21,14 +18,11 @@ use pyo3::{ use crate::{ coercion::FieldSpec, - errors::RustBridgeDeclined, http, python_settings::{PythonSettings, Snapshot}, + secrets, }; -const SECRET_MANAGER_READABLE: FieldSpec = - FieldSpec::new("readable", |field| field.schema_bool()); - const VERTEX_PROJECT: FieldSpec> = FieldSpec::new("vertex_project", |field| field.falsy_optional_string()); const VERTEX_LOCATION: FieldSpec> = @@ -58,7 +52,7 @@ fn run_ocr( kwargs: Bound<'_, PyDict>, asynchronous: bool, ) -> PyResult> { - let secrets = process_environment_secrets(&PythonSettings::SecretManager.read(py)?)?; + let secrets = secrets::source(py)?; let config = http::call_config(py, &kwargs, asynchronous)?; let client = OcrClient::new( http::pool(), @@ -79,15 +73,6 @@ fn run_ocr( ) } -fn process_environment_secrets(snapshot: &Snapshot<'_>) -> PyResult> { - if snapshot.read(&SECRET_MANAGER_READABLE)? { - return Err(RustBridgeDeclined::new_err( - "a readable secret manager is configured and the Rust route only reads the process environment", - )); - } - Ok(Arc::new(EnvironmentSecrets)) -} - fn ocr_settings(py: Python<'_>) -> PyResult { project_provider_defaults(&PythonSettings::ProviderDefaults.read(py)?) } @@ -123,38 +108,10 @@ pub(crate) fn aocr( #[cfg(test)] mod tests { - use pyo3::{prelude::*, types::PyDict}; - - use super::process_environment_secrets; - use crate::errors::RustBridgeDeclined; + use pyo3::prelude::*; use crate::python_settings::PythonSettings; - fn secret_manager<'py>(py: Python<'py>, readable: bool) -> Bound<'py, PyAny> { - let locals = PyDict::new(py); - locals.set_item("readable", readable).unwrap(); - py.run( - c"import types\nmanager = types.SimpleNamespace(readable=readable)", - Some(&locals), - Some(&locals), - ) - .unwrap(); - locals.get_item("manager").unwrap().unwrap() - } - - #[test] - fn a_readable_secret_manager_sends_the_call_back_to_python() { - Python::initialize(); - Python::attach(|py| { - let declined = process_environment_secrets( - &PythonSettings::SecretManager.snapshot(secret_manager(py, true)), - ) - .err() - .expect("the Rust route declines"); - assert!(declined.is_instance_of::(py)); - }); - } - #[test] fn provider_defaults_distinguish_falsey_values_and_exact_true() { Python::initialize(); diff --git a/litellm-rust/crates/python-bridge/src/secrets/callback.rs b/litellm-rust/crates/python-bridge/src/secrets/callback.rs index a8e5850fe3f..df9fc86f177 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/callback.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/callback.rs @@ -4,11 +4,17 @@ use litellm_core_utils::settings::Lookup; use litellm_secrets::{ Error, ExternalSecretManager, KeyManagementSettings, KeyManagementSystem, Secret, SecretValue, }; -use pyo3::{prelude::*, types::PyDict}; +use pyo3::{ + exceptions::PyException, + prelude::*, + types::{PyDict, PyString}, +}; -use super::error::external_error; +use super::error::{external_error, read_error}; const HANDLER_MODULE: &str = "litellm.secret_managers.secret_manager_handler"; +const ENVIRONMENT_FALLBACK_LOG: &str = + "Defaulting to os.environ value for key=%s. An exception occurred - %s.\n\n%s"; /// A secret manager whose reads execute in Python: a custom manager, a legacy compatible /// client, or a manually assigned SDK client. @@ -33,23 +39,6 @@ impl PythonSecretManager { fn read(&self, py: Python<'_>, name: &str) -> PyResult> { let client = self.client.bind(py); - if self.system == Some(KeyManagementSystem::Custom) - || (self.system.is_none() && client.hasattr("sync_read_secret")?) - { - let kwargs = PyDict::new(py); - kwargs.set_item("secret_name", name)?; - if self.system == Some(KeyManagementSystem::Custom) { - let optional_params = self - .settings - .as_ref() - .map(|settings| settings.bind(py).call_method0("model_dump")) - .transpose()?; - kwargs.set_item("optional_params", optional_params)?; - } - return client - .call_method("sync_read_secret", (), Some(&kwargs))? - .extract(); - } let kwargs = PyDict::new(py); kwargs.set_item("client", client)?; kwargs.set_item("key_manager", self.system.map_or("local", python_name))?; @@ -58,10 +47,15 @@ impl PythonSecretManager { Some(settings) => kwargs.set_item("key_management_settings", settings.bind(py))?, None => kwargs.set_item("key_management_settings", py.None())?, } - py.import(HANDLER_MODULE)? + let result = py + .import(HANDLER_MODULE)? .getattr("get_secret_from_manager")? - .call((), Some(&kwargs))? - .extract() + .call((), Some(&kwargs))?; + if result.is_instance_of::() { + result.extract().map(Some) + } else { + Ok(None) + } } } @@ -92,15 +86,35 @@ impl ExternalSecretManager for PythonSecretManager { _environment: &'a (dyn Lookup + Send + Sync), ) -> Pin, Error>> + Send + 'a>> { Box::pin(async move { - Python::attach(|py| { - self.read(py, name) - .map(|value| value.map(SecretValue::new).map(Secret::String)) - .map_err(|error| external_error(py, error)) + Python::attach(|py| match self.read(py, name) { + Ok(value) => Ok(value.map(SecretValue::new).map(Secret::String)), + // `get_secret` answers a failed manager read from the process environment, but + // only for `Exception`: cancellation and other `BaseException`s propagate. + Err(error) if error.is_instance_of::(py) => { + log_environment_fallback(py, name, &error) + .map_err(|error| external_error(py, error))?; + Err(read_error(py, error)) + } + Err(error) => Err(external_error(py, error)), }) }) } } +fn log_environment_fallback(py: Python<'_>, name: &str, error: &PyErr) -> PyResult<()> { + let traceback = py + .import("traceback")? + .call_method1("format_exception", (error.value(py),))?; + let traceback = "".into_pyobject(py)?.call_method1("join", (traceback,))?; + py.import("litellm._logging")? + .getattr("verbose_logger")? + .call_method1( + "error", + (ENVIRONMENT_FALLBACK_LOG, name, error.value(py), traceback), + )?; + Ok(()) +} + #[cfg(test)] mod tests { use std::sync::Arc; @@ -115,16 +129,12 @@ mod tests { use super::{HANDLER_MODULE, PythonSecretManager, python_name}; use crate::secrets::python_error; - #[rstest] - #[case::value_error("ValueError", None)] - #[case::value_error_with_fallback("ValueError", Some("environment-key"))] - #[case::cancelled("asyncio.CancelledError", None)] - #[case::cancelled_with_fallback("asyncio.CancelledError", Some("environment-key"))] - #[tokio::test] - async fn callback_failures_preserve_python_exceptions_even_with_environment_fallback( - #[case] failure_type: &str, - #[case] fallback: Option<&'static str>, - ) { + /// A resolver over a Python manager whose reads raise `failure_type`, with the chained + /// exceptions Python attaches, and `fallback` as the process environment. + fn failing_resolver( + failure_type: &str, + fallback: Option<&'static str>, + ) -> (SecretResolver, Py) { Python::initialize(); let (reader, locals) = Python::attach(|py| { let locals = PyDict::new(py); @@ -141,6 +151,13 @@ class Manager: def sync_read_secret(self, secret_name): raise failure manager = Manager() +import sys, types +for name in ('litellm', 'litellm.secret_managers'): + sys.modules.setdefault(name, types.ModuleType(name)) +handler = sys.modules.setdefault('litellm.secret_managers.secret_manager_handler', types.ModuleType('litellm.secret_managers.secret_manager_handler')) +def get_secret_from_manager(**kwargs): + return kwargs['client'].sync_read_secret(kwargs['secret_name']) +handler.get_secret_from_manager = get_secret_from_manager ", Some(&locals), Some(&locals), @@ -153,7 +170,7 @@ manager = Manager() ); (reader, locals.unbind()) }); - let resolver = SecretResolver::new( + let resolver = SecretResolver::new_python_compatible( Arc::new(SecretManagerState::new( SecretManager::External(Arc::new(reader)), KeyManagementSettings::default(), @@ -162,6 +179,19 @@ manager = Manager() OidcResolver::default(), ) .with_failure_policy(FailurePolicy::EnvironmentFallback); + (resolver, locals) + } + + #[rstest] + #[case::cancelled("asyncio.CancelledError", None)] + #[case::cancelled_with_fallback("asyncio.CancelledError", Some("environment-key"))] + #[case::keyboard_interrupt("KeyboardInterrupt", Some("environment-key"))] + #[tokio::test] + async fn base_exceptions_propagate_unchanged_even_with_environment_fallback( + #[case] failure_type: &str, + #[case] fallback: Option<&'static str>, + ) { + let (resolver, locals) = failing_resolver(failure_type, fallback); let error = resolver.get_secret("API_KEY", None).await.unwrap_err(); Python::attach(|py| { let original = python_error(py, &error).unwrap(); @@ -184,24 +214,96 @@ manager = Manager() }); } + /// Installs a persistent `litellm._logging` stub whose `verbose_logger.error` records its + /// arguments, and returns those recorded for `name`. + fn logged_errors<'py>(py: Python<'py>, name: &str) -> Vec> { + py.run( + c" +import sys, types +class Logger: + calls = [] + def error(self, *args): + self.calls.append(args) +logging = types.ModuleType('litellm._logging') +logging.verbose_logger = Logger() +sys.modules.setdefault('litellm', types.ModuleType('litellm')) +sys.modules.setdefault('litellm._logging', logging) +", + None, + None, + ) + .unwrap(); + py.import("litellm._logging") + .unwrap() + .getattr("verbose_logger") + .unwrap() + .getattr("calls") + .unwrap() + .try_iter() + .unwrap() + .map(Result::unwrap) + .filter(|call| call.get_item(1).unwrap().extract::().unwrap() == name) + .collect() + } + + #[rstest] + #[case::value_error("ValueError", None, "FALLBACK_VALUE_ERROR")] + #[case::value_error_with_fallback( + "ValueError", + Some("environment-key"), + "FALLBACK_VALUE_ERROR_WITH_ENVIRONMENT" + )] + #[case::runtime_error_with_fallback( + "RuntimeError", + Some("environment-key"), + "FALLBACK_RUNTIME_ERROR_WITH_ENVIRONMENT" + )] + #[tokio::test] + async fn exceptions_are_logged_and_answered_from_the_environment( + #[case] failure_type: &str, + #[case] fallback: Option<&'static str>, + #[case] name: &str, + ) { + let (resolver, _locals) = failing_resolver(failure_type, fallback); + Python::attach(|py| assert!(logged_errors(py, name).is_empty())); + let secret = resolver.get_secret(name, None).await.unwrap(); + assert_eq!( + secret.map(|secret| match secret { + litellm_secrets::Secret::String(value) => value.expose().to_owned(), + other => panic!("unexpected secret {other:?}"), + }), + fallback.map(str::to_owned) + ); + Python::attach(|py| { + let calls = logged_errors(py, name); + assert_eq!(calls.len(), 1); + assert!( + calls[0] + .get_item(3) + .unwrap() + .extract::() + .unwrap() + .contains("sync_read_secret") + ); + }); + } + /// Installs a fake `get_secret_from_manager` that records its kwargs, runs `body`, and - /// removes the fake modules again. + /// removes the fake handler again; parent package stubs persist for concurrent tests. fn with_fake_handler<'py>(py: Python<'py>, body: impl FnOnce(&Bound<'py, PyDict>)) { let locals = PyDict::new(py); py.run( c" import sys, types +previous_handler = sys.modules.get('litellm.secret_managers.secret_manager_handler') calls = [] def get_secret_from_manager(**kwargs): calls.append(kwargs) return 'handled-' + kwargs['secret_name'] handler = types.ModuleType('litellm.secret_managers.secret_manager_handler') handler.get_secret_from_manager = get_secret_from_manager -installed = {} for name in ('litellm', 'litellm.secret_managers'): - if name not in sys.modules: - sys.modules[name] = types.ModuleType(name) - installed[name] = True + sys.modules.setdefault(name, types.ModuleType(name)) sys.modules['litellm.secret_managers.secret_manager_handler'] = handler ", Some(&locals), @@ -211,9 +313,10 @@ sys.modules['litellm.secret_managers.secret_manager_handler'] = handler body(&locals); py.run( c" -sys.modules.pop('litellm.secret_managers.secret_manager_handler', None) -for name in installed: - sys.modules.pop(name, None) +if previous_handler is None: + sys.modules.pop('litellm.secret_managers.secret_manager_handler', None) +else: + sys.modules['litellm.secret_managers.secret_manager_handler'] = previous_handler ", Some(&locals), Some(&locals), @@ -221,6 +324,28 @@ for name in installed: .unwrap(); } + #[rstest] + #[case("None")] + #[case("True")] + #[case("123")] + #[case("{'key': 'value'}")] + fn nonstring_results_are_absent_without_a_read_failure(#[case] expression: &str) { + Python::initialize(); + Python::attach(|py| { + with_fake_handler(py, |locals| { + locals.set_item("expression", expression).unwrap(); + py.run( + c"handler.get_secret_from_manager = lambda **kwargs: eval(expression)", + Some(locals), + Some(locals), + ) + .unwrap(); + let reader = PythonSecretManager::new(py.None(), None, None); + assert_eq!(reader.read(py, "KEY").unwrap(), None); + }); + }); + } + #[rstest] #[case::google_kms(KeyManagementSystem::GoogleKms)] #[case::azure_key_vault(KeyManagementSystem::AzureKeyVault)] @@ -238,72 +363,6 @@ for name in installed: ); } - #[rstest] - #[case::legacy(None, false)] - #[case::custom(Some(KeyManagementSystem::Custom), true)] - fn direct_readers_receive_compatible_kwargs( - #[case] system: Option, - #[case] expects_optional_params: bool, - ) { - Python::initialize(); - Python::attach(|py| { - let locals = PyDict::new(py); - py.run( - c" -class Settings: - def model_dump(self): - return {'scope': 'custom'} -class Manager: - def __init__(self): - self.names = [] - self.optional_params = [] - def sync_read_secret(self, secret_name, optional_params=None, timeout=None): - self.names.append(secret_name) - self.optional_params.append(optional_params) - return 'direct-' + secret_name -manager = Manager() -settings = Settings() -", - Some(&locals), - Some(&locals), - ) - .unwrap(); - let manager = locals.get_item("manager").unwrap().unwrap(); - let settings = expects_optional_params - .then(|| locals.get_item("settings").unwrap().unwrap().unbind()); - let reader = PythonSecretManager::new(manager.clone().unbind(), system, settings); - assert_eq!( - reader.read(py, "API_KEY").unwrap().as_deref(), - Some("direct-API_KEY") - ); - assert_eq!( - manager - .getattr("names") - .unwrap() - .extract::>() - .unwrap(), - ["API_KEY"] - ); - let optional_params = manager - .getattr("optional_params") - .unwrap() - .get_item(0) - .unwrap(); - if expects_optional_params { - assert_eq!( - optional_params - .get_item("scope") - .unwrap() - .extract::() - .unwrap(), - "custom" - ); - } else { - assert!(optional_params.is_none()); - } - }); - } - #[test] fn configured_systems_dispatch_through_the_python_handler_with_the_original_settings() { Python::initialize(); @@ -349,4 +408,57 @@ settings = Settings() }); }); } + + #[rstest] + #[case::manually_assigned(None, "local")] + #[case::custom(Some(KeyManagementSystem::Custom), "custom")] + fn direct_readers_dispatch_through_the_python_handler_like_get_secret( + #[case] system: Option, + #[case] key_manager: &str, + ) { + Python::initialize(); + Python::attach(|py| { + with_fake_handler(py, |locals| { + py.run( + c" +class Manager: + def __init__(self): + self.names = [] + def sync_read_secret(self, secret_name, optional_params=None, timeout=None): + self.names.append(secret_name) + return 'direct-' + secret_name +manager = Manager() +", + Some(locals), + Some(locals), + ) + .unwrap(); + let manager = locals.get_item("manager").unwrap().unwrap(); + let reader = PythonSecretManager::new(manager.clone().unbind(), system, None); + assert_eq!( + reader.read(py, "API_KEY").unwrap().as_deref(), + Some("handled-API_KEY") + ); + assert_eq!( + manager + .getattr("names") + .unwrap() + .extract::>() + .unwrap(), + Vec::::new() + ); + let calls = locals.get_item("calls").unwrap().unwrap(); + let call = calls.get_item(0).unwrap().cast_into::().unwrap(); + assert!(call.get_item("client").unwrap().unwrap().is(&manager)); + assert_eq!( + call.get_item("key_manager") + .unwrap() + .unwrap() + .extract::() + .unwrap(), + key_manager + ); + }); + }); + } } diff --git a/litellm-rust/crates/python-bridge/src/secrets/config.rs b/litellm-rust/crates/python-bridge/src/secrets/config.rs index 6fd380fe40c..4a666ca1ded 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/config.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/config.rs @@ -56,13 +56,23 @@ const SETTINGS_OBJECT: FieldSpec>> = FieldSpec::new("settings_object", |field| Ok(field.python_binding())); /// `litellm.secret_manager_client` as the bridge classifies it. -#[derive(Debug)] pub(crate) enum SecretManagerClient { /// `None`: reads come from the process environment. Local, /// A custom manager, legacy compatible client, or manually assigned SDK client that keeps /// executing in Python. PythonCallback(Py), + Native(Box), +} + +impl std::fmt::Debug for SecretManagerClient { + fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result { + formatter.write_str(match self { + Self::Local => "Local", + Self::Native(_) => "Native", + Self::PythonCallback(_) => "PythonCallback", + }) + } } /// One operation-local capture of the secret manager globals, taken while attached to Python. @@ -79,6 +89,9 @@ pub(crate) struct SecretManagerSnapshot { impl SecretManagerSnapshot { pub(crate) fn into_state(self) -> Arc { match self.client { + SecretManagerClient::Native(backend) => { + Arc::new(SecretManagerState::new(*backend, self.settings)) + } SecretManagerClient::Local => Arc::new(SecretManagerState::default()), SecretManagerClient::PythonCallback(client) => Arc::new(SecretManagerState::new( SecretManager::External(Arc::new(PythonSecretManager::new( @@ -94,7 +107,32 @@ impl SecretManagerSnapshot { /// Reads and projects the secret manager settings group in one attached operation. pub(crate) fn read(py: Python<'_>) -> PyResult { - Ok(project(&PythonSettings::SecretManagerBinding.read(py)?)?) + let snapshot = project(&PythonSettings::SecretManagerBinding.read(py)?)?; + let SecretManagerClient::PythonCallback(client) = &snapshot.client else { + return Ok(snapshot); + }; + if matches!( + snapshot.system, + Some(KeyManagementSystem::Custom | KeyManagementSystem::Local) + ) { + return Ok(snapshot); + } + let Some(native) = super::runtime::NativeSecretManager::from_client(client.bind(py))? else { + return Ok(snapshot); + }; + let backend = native.borrow(py).backend()?; + if snapshot + .system + .is_some_and(|system| system != backend.system()) + { + return Err(pyo3::exceptions::PyValueError::new_err( + "native secret manager system does not match configuration", + )); + } + Ok(SecretManagerSnapshot { + client: SecretManagerClient::Native(Box::new(backend)), + ..snapshot + }) } pub(crate) fn project(snapshot: &Snapshot<'_>) -> Result { diff --git a/litellm-rust/crates/python-bridge/src/secrets/error.rs b/litellm-rust/crates/python-bridge/src/secrets/error.rs index 1bf9351ee64..5ccfd0a1554 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/error.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/error.rs @@ -17,8 +17,12 @@ pub(super) fn external_error(py: Python<'_>, error: PyErr) -> Error { Error::ExternalManager(Box::new(PythonSecretError(error.into_value(py)))) } +pub(super) fn read_error(py: Python<'_>, error: PyErr) -> Error { + Error::ExternalRead(Box::new(PythonSecretError(error.into_value(py)))) +} + pub(crate) fn python_error(py: Python<'_>, error: &Error) -> Option { - let Error::ExternalManager(source) = error else { + let (Error::ExternalManager(source) | Error::ExternalRead(source)) = error else { return None; }; source diff --git a/litellm-rust/crates/python-bridge/src/secrets/mod.rs b/litellm-rust/crates/python-bridge/src/secrets/mod.rs index c753015aa9e..f0ad98dbbd8 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/mod.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/mod.rs @@ -1,6 +1,105 @@ pub(crate) mod callback; pub(crate) mod config; mod error; +mod mutation; +mod operations; +mod provider; pub(crate) mod resolved; +pub(crate) mod runtime; +mod vault; + +use std::sync::Arc; + +use litellm_secrets::source::{EnvironmentSecrets, SecretSource}; +use pyo3::prelude::*; pub(crate) use error::python_error; +use resolved::ResolvedSecrets; + +use crate::{ + coercion::FieldSpec, + errors::RustBridgeDeclined, + python_settings::{PythonSettings, Snapshot}, +}; + +const READABLE: FieldSpec = FieldSpec::new("readable", |field| field.schema_bool()); +const NATIVE: FieldSpec = FieldSpec::new("native", |field| field.schema_bool()); + +/// Where a Rust route reads provider secrets from, as `litellm.get_secret` would. +pub(crate) fn source(py: Python<'_>) -> PyResult> { + select(&PythonSettings::SecretManager.read(py)?, || { + Ok(Arc::new(ResolvedSecrets::new(config::read(py)?))) + }) +} + +fn select( + manager: &Snapshot<'_>, + resolved: impl FnOnce() -> PyResult>, +) -> PyResult> { + if !manager.read(&READABLE)? { + return Ok(Arc::new(EnvironmentSecrets::python_compatible())); + } + if !manager.read(&NATIVE)? { + return Err(RustBridgeDeclined::new_err( + "the configured secret manager is not enabled for the Rust bridge", + )); + } + resolved() +} + +#[cfg(test)] +mod tests { + use std::sync::Arc; + + use litellm_secrets::source::{EnvironmentSecrets, SecretSource}; + use pyo3::{prelude::*, types::PyDict}; + use rstest::rstest; + + use super::select; + use crate::{errors::RustBridgeDeclined, python_settings::PythonSettings}; + + enum Selected { + Environment, + Declined, + Resolved, + } + + #[rstest] + #[case::unreadable(false, false, Selected::Environment)] + #[case::unreadable_even_if_native(false, true, Selected::Environment)] + #[case::readable_python_only(true, false, Selected::Declined)] + #[case::readable_native(true, true, Selected::Resolved)] + fn readable_and_native_select_the_secret_source( + #[case] readable: bool, + #[case] native: bool, + #[case] expected: Selected, + ) { + Python::initialize(); + Python::attach(|py| { + let locals = PyDict::new(py); + locals.set_item("readable", readable).unwrap(); + locals.set_item("native", native).unwrap(); + let manager = py + .eval( + c"__import__('types').SimpleNamespace(readable=readable, native=native)", + None, + Some(&locals), + ) + .unwrap(); + let mut resolved_called = false; + let selected = select(&PythonSettings::SecretManager.snapshot(manager), || { + resolved_called = true; + Ok(Arc::new(EnvironmentSecrets::python_compatible()) as Arc) + }); + match expected { + Selected::Environment => assert!(selected.is_ok() && !resolved_called), + Selected::Resolved => assert!(selected.is_ok() && resolved_called), + Selected::Declined => { + let error = selected.err().expect("the Rust route declines"); + assert!(error.is_instance_of::(py)); + assert!(!resolved_called); + } + } + }); + } +} diff --git a/litellm-rust/crates/python-bridge/src/secrets/mutation.rs b/litellm-rust/crates/python-bridge/src/secrets/mutation.rs new file mode 100644 index 00000000000..eaab06a5b12 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/secrets/mutation.rs @@ -0,0 +1,113 @@ +use super::operations::{PythonMutationError, PythonMutationResponse}; +use litellm_host_python::{json_loads, to_py}; +use litellm_secrets::cyberark; +use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict}; + +pub(super) fn mutation_value( + result: Result, + context: &super::vault::ErrorContext, +) -> PyResult> { + Python::attach(|py| match result { + Ok(PythonMutationResponse::Value(value)) => to_py(py, &value), + Ok(PythonMutationResponse::Json(body)) => match json_value(py, &body) { + Ok(value) => Ok(value), + Err(error) => error_value(py, error.value(py).str()?.extract()?), + }, + Err(PythonMutationError::Vault(failure)) => { + super::vault::failure_value(py, *failure, context) + } + Err(PythonMutationError::CyberarkWrite { name, failure }) => { + let message = cyberark_failure(py, &name, *failure)?; + to_py( + py, + &serde_json::json!({"status": "error", "message": message}), + ) + } + Err(PythonMutationError::CurrentMissing(name)) => Err(PyValueError::new_err(format!( + "Current secret {name} not found" + ))), + Err(PythonMutationError::ReplacementMissing(name)) => Err(PyValueError::new_err(format!( + "Failed to verify new secret {name}" + ))), + Err(PythonMutationError::ReplacementMismatch) => { + Err(PyValueError::new_err("New secret value mismatch")) + } + Err(PythonMutationError::Unsupported) => Err(PyValueError::new_err( + "native secret manager mutation is unavailable", + )), + }) +} + +fn cyberark_failure( + py: Python<'_>, + name: &str, + failure: cyberark::WriteFailure, +) -> PyResult { + let message = match failure.source { + cyberark::Error::Status(status) | cyberark::Error::AuthStatus(status) => { + let url = failure + .request_url + .as_ref() + .map_or("", reqwest::Url::as_str); + http_message(py, "POST", url, status)? + } + cyberark::Error::Operation(litellm_secrets_types::Error::UnsafeSecretName) => { + format!("Invalid secret_name {}", name.into_pyobject(py)?.repr()?) + } + cyberark::Error::Http(source) if failure.authentication => match os_error_code(&source) { + Some(code) => { + let reason = py.import("os")?.getattr("strerror")?.call1((code,))?; + py.import("builtins")? + .getattr("OSError")? + .call1((code, reason))? + .str()? + .extract()? + } + None => cyberark::Error::Http(source).to_string(), + }, + cyberark::Error::Http(source) if source.is_connect() => { + "All connection attempts failed".to_owned() + } + source => source.to_string(), + }; + Ok(if failure.authentication { + format!("Could not authenticate to CyberArk Conjur: {message}") + } else { + message + }) +} + +fn os_error_code(error: &(dyn std::error::Error + 'static)) -> Option { + error + .downcast_ref::() + .and_then(std::io::Error::raw_os_error) + .or_else(|| error.source().and_then(os_error_code)) +} + +pub(super) fn json_value(py: Python<'_>, body: &[u8]) -> PyResult> { + json_loads(py, body) +} + +pub(super) fn error_value(py: Python<'_>, message: String) -> PyResult> { + to_py( + py, + &serde_json::json!({"status": "error", "message": message}), + ) +} + +pub(super) fn http_message( + py: Python<'_>, + method: &str, + url: &str, + status: u16, +) -> PyResult { + let httpx = py.import("httpx")?; + let request = httpx.getattr("Request")?.call1((method, url))?; + let kwargs = PyDict::new(py); + kwargs.set_item("request", request)?; + let response = httpx.getattr("Response")?.call((status,), Some(&kwargs))?; + match response.call_method0("raise_for_status") { + Err(error) => error.value(py).str()?.extract(), + Ok(_) => Ok(format!("HTTP {status}")), + } +} diff --git a/litellm-rust/crates/python-bridge/src/secrets/operations.rs b/litellm-rust/crates/python-bridge/src/secrets/operations.rs new file mode 100644 index 00000000000..8e338741aff --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/secrets/operations.rs @@ -0,0 +1,247 @@ +use litellm_core_utils::settings::Lookup; +use litellm_secrets::Secret; +use litellm_secrets::cyberark::AuthenticationRetry; +use litellm_secrets::{Error, SecretManager}; +use litellm_secrets_types::{PythonSecretRead, SecretOperationContext}; + +pub(super) struct PythonReadRequest { + pub secret_name: String, + pub primary_secret_name: Option, + pub context: SecretOperationContext, + pub synchronous: bool, +} + +pub(super) async fn read_python_provider( + manager: &SecretManager, + request: &PythonReadRequest, + _environment: &(dyn Lookup + Send + Sync), +) -> Result { + match (manager, &request.context) { + (SecretManager::AwsSecretsManagerV2(client), SecretOperationContext::Aws(context)) => { + client + .read_provider_payload_for_python( + &request.secret_name, + request.primary_secret_name.as_deref(), + context, + request.synchronous, + _environment, + ) + .await + .map_err(Error::from) + } + (SecretManager::HashicorpVault(client), SecretOperationContext::Hashicorp(context)) => { + Ok(PythonSecretRead::Value( + client + .async_read_secret_with_context(&request.secret_name, context) + .await + .unwrap_or(None) + .map(Secret::String), + )) + } + (SecretManager::Cyberark(client), SecretOperationContext::Cyberark(_)) => { + Ok(PythonSecretRead::Value( + client + .read_with_retry( + &request.secret_name, + &Default::default(), + AuthenticationRetry::Never, + ) + .await + .unwrap_or(None) + .map(Secret::String), + )) + } + (SecretManager::GoogleSecretManager(client), SecretOperationContext::Google(_)) => client + .get_secret_for_python(&request.secret_name) + .await + .map(PythonSecretRead::Value) + .map_err(Error::from), + _ => Err(Error::NativeBackendUnavailable), + } +} + +#[derive(Debug)] +pub(super) enum PythonMutationError { + Unsupported, + Vault(Box), + CyberarkWrite { + name: String, + failure: Box, + }, + CurrentMissing(String), + ReplacementMissing(String), + ReplacementMismatch, +} + +pub(super) async fn write_python_provider( + manager: &SecretManager, + name: &str, + value: &litellm_secrets::SecretValue, +) -> Result { + match manager { + SecretManager::Cyberark(client) => { + client + .write_with_retry(name, value, &Default::default(), AuthenticationRetry::Never) + .await + .map_err(|failure| PythonMutationError::CyberarkWrite { + name: name.to_owned(), + failure: Box::new(failure), + })?; + Ok(write_success(name)) + } + _ => Err(PythonMutationError::Unsupported), + } +} + +pub(super) async fn delete_python_provider( + manager: &SecretManager, + name: &str, +) -> Result { + match manager { + SecretManager::Cyberark(client) => { + client + .async_delete_secret(name, None) + .await + .map_err(|failure| PythonMutationError::CyberarkWrite { + name: name.to_owned(), + failure: Box::new(litellm_secrets::cyberark::WriteFailure { + source: failure, + request_url: None, + authentication: false, + }), + })?; + Ok(serde_json::json!({ + "status": "not_supported", + "message": "CyberArk Conjur does not support direct secret deletion. Use policy updates to remove variables.", + })) + } + _ => Err(PythonMutationError::Unsupported), + } +} + +pub(super) async fn rotate_python_provider( + manager: &SecretManager, + current_name: &str, + new_name: &str, + value: &litellm_secrets::SecretValue, +) -> Result { + match manager { + SecretManager::Cyberark(client) => { + if client + .read_fresh_with_retry( + current_name, + &Default::default(), + AuthenticationRetry::Never, + ) + .await + .ok() + .flatten() + .is_none() + { + return Err(PythonMutationError::CurrentMissing(current_name.to_owned())); + } + client + .write_with_retry( + new_name, + value, + &Default::default(), + AuthenticationRetry::Never, + ) + .await + .map_err(|failure| PythonMutationError::CyberarkWrite { + name: new_name.to_owned(), + failure: Box::new(failure), + })?; + let actual = client + .read_fresh_with_retry(new_name, &Default::default(), AuthenticationRetry::Never) + .await + .ok() + .flatten() + .ok_or_else(|| PythonMutationError::ReplacementMissing(new_name.to_owned()))?; + if actual != *value { + return Err(PythonMutationError::ReplacementMismatch); + } + if current_name != new_name { + client.invalidate_cached_secret(current_name).await; + } + Ok(write_success(new_name)) + } + _ => Err(PythonMutationError::Unsupported), + } +} +fn write_success(name: &str) -> serde_json::Value { + serde_json::json!({"status": "success", "message": format!("Secret {name} written successfully")}) +} + +pub(super) enum PythonMutationResponse { + Value(serde_json::Value), + Json(Vec), +} + +pub(super) async fn write_python_provider_with_context( + manager: &SecretManager, + name: &str, + value: &litellm_secrets::SecretValue, + context: &litellm_secrets_types::SecretWriteContext, +) -> Result { + if let (SecretManager::HashicorpVault(client), SecretOperationContext::Hashicorp(operation)) = + (manager, &context.operation) + { + return super::vault::write( + client, + name, + value, + &litellm_secrets_types::SecretWriteContext { + description: context.description.clone(), + tags: context.tags.clone(), + operation: operation.clone(), + }, + ) + .await + .map(PythonMutationResponse::Json) + .map_err(|failure| PythonMutationError::Vault(Box::new(failure))); + } + write_python_provider(manager, name, value) + .await + .map(PythonMutationResponse::Value) +} + +pub(super) async fn delete_python_provider_with_context( + manager: &SecretManager, + name: &str, + context: &SecretOperationContext, +) -> Result { + if let (SecretManager::HashicorpVault(client), SecretOperationContext::Hashicorp(context)) = + (manager, context) + { + super::vault::delete(client, name, context) + .await + .map_err(|failure| PythonMutationError::Vault(Box::new(failure)))?; + return Ok(PythonMutationResponse::Value(serde_json::json!({ + "status": "success", "message": format!("Secret {name} deleted successfully"), + }))); + } + delete_python_provider(manager, name) + .await + .map(PythonMutationResponse::Value) +} + +pub(super) async fn rotate_python_provider_with_context( + manager: &SecretManager, + current_name: &str, + new_name: &str, + value: &litellm_secrets::SecretValue, + context: &SecretOperationContext, +) -> Result { + if let (SecretManager::HashicorpVault(client), SecretOperationContext::Hashicorp(context)) = + (manager, context) + { + return super::vault::rotate(client, current_name, new_name, value, context) + .await + .map(PythonMutationResponse::Json) + .map_err(|failure| PythonMutationError::Vault(Box::new(failure))); + } + rotate_python_provider(manager, current_name, new_name, value) + .await + .map(PythonMutationResponse::Value) +} diff --git a/litellm-rust/crates/python-bridge/src/secrets/provider.rs b/litellm-rust/crates/python-bridge/src/secrets/provider.rs new file mode 100644 index 00000000000..568a0cd2228 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/secrets/provider.rs @@ -0,0 +1,150 @@ +use std::time::Duration; + +use super::operations::PythonReadRequest; +use litellm_secrets::{KeyManagementSystem, SecretValue}; +use litellm_secrets_types::{ + AwsOperationContext, CyberarkOperationContext, GoogleOperationContext, + HashicorpOperationContext, SecretOperationContext, +}; +use pyo3::{exceptions::PyValueError, prelude::*, types::PyDict}; + +pub(super) fn read_request( + system: KeyManagementSystem, + secret_name: String, + optional_params: Option<&Bound<'_, PyAny>>, + timeout: Option<&Bound<'_, PyAny>>, + primary_secret_name: Option, + synchronous: bool, +) -> PyResult { + let context = match system { + KeyManagementSystem::AwsSecretManager => { + let ignored = primary_secret_name + .as_ref() + .is_some_and(|value| !value.is_empty()) + || (synchronous + && litellm_secrets::aws::secret_manager::is_bootstrap_key(&secret_name)); + SecretOperationContext::Aws(if ignored { + AwsOperationContext::default() + } else { + aws_context(optional_params, timeout)? + }) + } + KeyManagementSystem::HashicorpVault => { + SecretOperationContext::Hashicorp(vault_context(optional_params)?) + } + KeyManagementSystem::Cyberark => { + SecretOperationContext::Cyberark(CyberarkOperationContext::default()) + } + KeyManagementSystem::GoogleSecretManager => { + SecretOperationContext::Google(GoogleOperationContext::default()) + } + _ => { + return Err(PyValueError::new_err( + "secret manager does not support provider reads", + )); + } + }; + Ok(PythonReadRequest { + secret_name, + primary_secret_name, + context, + synchronous, + }) +} + +fn string_field(params: Option<&Bound<'_, PyDict>>, name: &str) -> PyResult> { + let value = params + .map(|params| params.get_item(name)) + .transpose()? + .flatten(); + match value { + Some(value) if value.is_truthy()? => value.extract().map(Some), + _ => Ok(None), + } +} + +fn aws_context( + params: Option<&Bound<'_, PyAny>>, + timeout: Option<&Bound<'_, PyAny>>, +) -> PyResult { + let params = params + .filter(|value| !value.is_none()) + .map(|value| value.cast::()) + .transpose()?; + Ok(AwsOperationContext { + access_key_id: string_field(params, "aws_access_key_id")?.map(SecretValue::new), + secret_access_key: string_field(params, "aws_secret_access_key")?.map(SecretValue::new), + session_token: string_field(params, "aws_session_token")?.map(SecretValue::new), + region_name: string_field(params, "aws_region_name")?, + role_name: string_field(params, "aws_role_name")?, + session_name: string_field(params, "aws_session_name")?, + external_id: string_field(params, "aws_external_id")?.map(SecretValue::new), + profile_name: string_field(params, "aws_profile_name")?, + web_identity_token: string_field(params, "aws_web_identity_token")?.map(SecretValue::new), + sts_endpoint: string_field(params, "aws_sts_endpoint")?, + bedrock_runtime_endpoint: string_field(params, "aws_bedrock_runtime_endpoint")?, + timeout: read_timeout(timeout)?, + }) +} + +fn read_timeout(value: Option<&Bound<'_, PyAny>>) -> PyResult> { + let Some(value) = value.filter(|value| !value.is_none()) else { + return Ok(None); + }; + let seconds = match value.extract::() { + Ok(value) => Some(value), + Err(_) => value.getattr("read")?.extract::>()?, + }; + seconds + .map(|value| { + Duration::try_from_secs_f64(value).map_err(|_| PyValueError::new_err("invalid timeout")) + }) + .transpose() +} + +fn vault_context(params: Option<&Bound<'_, PyAny>>) -> PyResult { + let params = params.and_then(|value| value.cast::().ok()); + let nested = params + .map(|params| params.get_item("secret_manager_settings")) + .transpose()? + .flatten(); + let source = nested + .as_ref() + .and_then(|value| value.cast::().ok()) + .or(params); + Ok(HashicorpOperationContext { + namespace: vault_field(source, "namespace")?, + mount: vault_field(source, "mount")?, + path_prefix: vault_field(source, "path_prefix")?, + data_key: vault_field(source, "data")?, + timeout: None, + }) +} + +fn vault_field(params: Option<&Bound<'_, PyDict>>, name: &str) -> PyResult> { + let value = params + .map(|params| params.get_item(name)) + .transpose()? + .flatten(); + match value { + Some(value) if value.is_none() => Ok(None), + Some(value) => value.str()?.extract().map(Some), + None => Ok(None), + } +} + +pub(super) fn mutation_context( + system: KeyManagementSystem, + optional_params: Option<&Bound<'_, PyAny>>, + timeout: Option<&Bound<'_, PyAny>>, +) -> PyResult { + if system == KeyManagementSystem::HashicorpVault { + return Ok(SecretOperationContext::Hashicorp( + HashicorpOperationContext { + timeout: read_timeout(timeout)?, + ..vault_context(optional_params)? + }, + )); + } + Ok(SecretOperationContext::Default) +} diff --git a/litellm-rust/crates/python-bridge/src/secrets/resolved.rs b/litellm-rust/crates/python-bridge/src/secrets/resolved.rs index 877c429169a..84594c5a3dd 100644 --- a/litellm-rust/crates/python-bridge/src/secrets/resolved.rs +++ b/litellm-rust/crates/python-bridge/src/secrets/resolved.rs @@ -1,10 +1,10 @@ -use std::{collections::HashMap, sync::Arc}; +use std::sync::Arc; -use futures_util::{future::BoxFuture, future::try_join_all}; -use litellm_core_utils::settings::{Lookup, ProcessEnvironment}; -use litellm_llms::base_llm::inference::secrets::{SecretSource, Secrets}; +use futures_util::future::BoxFuture; +use litellm_core_utils::settings::ProcessEnvironment; +use litellm_secrets::source::SecretSource; use litellm_secrets::{ - Error, FailurePolicy, OidcResolver, Secret, SecretManagerState, SecretResolver, + Error, FailurePolicy, OidcResolver, SecretManagerState, SecretResolver, SecretValue, }; use super::config::SecretManagerSnapshot; @@ -20,7 +20,7 @@ impl ResolvedSecrets { fn from_state(state: Arc) -> Self { Self { - resolver: SecretResolver::new( + resolver: SecretResolver::new_python_compatible( state, Arc::new(ProcessEnvironment), OidcResolver::default(), @@ -31,41 +31,11 @@ impl ResolvedSecrets { } impl SecretSource for ResolvedSecrets { - fn resolve<'a>(&'a self, names: &'a [&'static str]) -> BoxFuture<'a, Result> { - Box::pin(async move { - let values = try_join_all(names.iter().map(|name| async move { - self.resolver - .get_secret(name, None) - .await - .map(|secret| secret.map(|secret| ((*name).to_owned(), secret_value(secret)))) - })) - .await? - .into_iter() - .flatten() - .collect::>(); - Ok(Arc::new(ResolvedLookup { values }) as Secrets) - }) - } -} - -struct ResolvedLookup { - values: HashMap, -} - -impl Lookup for ResolvedLookup { - fn get(&self, name: &str) -> Option { - self.values - .get(name) - .cloned() - .or_else(|| ProcessEnvironment.get(name)) - } -} - -fn secret_value(secret: Secret) -> String { - match secret { - Secret::String(value) => value.expose().to_owned(), - Secret::Bool(value) => if value { "True" } else { "False" }.to_owned(), - Secret::Json(value) => value.to_string(), + fn get_secret_str<'a>( + &'a self, + name: &'a str, + ) -> BoxFuture<'a, Result, Error>> { + Box::pin(self.resolver.get_secret_str(name, None)) } } @@ -86,7 +56,7 @@ mod tests { }; use super::ResolvedSecrets; - use litellm_llms::base_llm::inference::secrets::SecretSource; + use litellm_secrets::source::SecretSource; fn state(server: &MockServer, settings: KeyManagementSettings) -> Arc { let client = Client::from_conf( @@ -145,31 +115,75 @@ mod tests { } #[tokio::test] - async fn manager_failure_falls_back_to_environment() { + async fn aws_read_failure_preserves_absence_without_environment_fallback() { let name = "LITELLM_RUST_BRIDGE_MANAGER_FAILURE"; unsafe { std::env::set_var(name, "env-key") }; let server = MockServer::start().await; Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) .respond_with(ResponseTemplate::new(500)) - .expect(1) + .expect(2) .mount(&server) .await; let result = resolve(state(&server, KeyManagementSettings::default()), name).await; + let missing = resolve( + state(&server, KeyManagementSettings::default()), + "LITELLM_RUST_BRIDGE_MANAGER_FAILURE_MISSING", + ) + .await; unsafe { std::env::remove_var(name) }; - assert_eq!(result.as_deref(), Some("env-key")); - assert_eq!(server.received_requests().await.unwrap().len(), 1); + assert_eq!(result, None); + assert_eq!(missing, None); + } - let missing_server = MockServer::start().await; + #[rstest::rstest] + #[case::capitalized_true("True")] + #[case::parenthesized_false("(False)")] + #[tokio::test] + async fn boolean_manager_values_are_absent_like_get_secret_str(#[case] value: &str) { + let name = "LITELLM_RUST_BRIDGE_BOOLEAN_VALUE"; + let server = MockServer::start().await; Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) - .respond_with(ResponseTemplate::new(500)) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"SecretString": value}))) .expect(1) - .mount(&missing_server) + .mount(&server) .await; - let missing = - ResolvedSecrets::from_state(state(&missing_server, KeyManagementSettings::default())) - .resolve(&["LITELLM_RUST_BRIDGE_MANAGER_FAILURE_MISSING"]) - .await; - assert!(matches!(missing, Err(litellm_secrets::Error::Aws(_)))); + assert_eq!( + resolve(state(&server, KeyManagementSettings::default()), name).await, + None + ); + } + + #[tokio::test] + async fn undeclared_names_are_read_from_the_manager() { + let declared = "LITELLM_RUST_BRIDGE_DECLARED"; + let undeclared = "LITELLM_RUST_BRIDGE_UNDECLARED_MANAGED"; + unsafe { std::env::set_var(undeclared, "env-key") }; + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) + .and(body_partial_json(json!({"SecretId": declared}))) + .respond_with( + ResponseTemplate::new(200).set_body_json(json!({"SecretString": "declared-key"})), + ) + .mount(&server) + .await; + Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) + .and(body_partial_json(json!({"SecretId": undeclared}))) + .respond_with( + ResponseTemplate::new(200).set_body_json(json!({"SecretString": "manager-key"})), + ) + .expect(1) + .mount(&server) + .await; + let source = ResolvedSecrets::from_state(state(&server, KeyManagementSettings::default())); + let snapshot = source.resolve(&[declared]).await.unwrap(); + assert_eq!(snapshot.get(undeclared), None); + let result = source + .get_secret_str(undeclared) + .await + .unwrap() + .map(|value| value.expose().to_owned()); + unsafe { std::env::remove_var(undeclared) }; + assert_eq!(result.as_deref(), Some("manager-key")); } #[tokio::test] @@ -230,7 +244,7 @@ mod tests { } #[tokio::test] - async fn undeclared_names_still_read_the_process_environment() { + async fn names_excluded_by_hosted_keys_read_the_process_environment() { let name = "LITELLM_RUST_BRIDGE_UNDECLARED"; unsafe { std::env::set_var(name, "env-key") }; let server = MockServer::start().await; diff --git a/litellm-rust/crates/python-bridge/src/secrets/runtime.rs b/litellm-rust/crates/python-bridge/src/secrets/runtime.rs new file mode 100644 index 00000000000..1a89130ee82 --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/secrets/runtime.rs @@ -0,0 +1,463 @@ +use std::{collections::BTreeMap, sync::Arc}; + +use litellm_core_utils::settings::{Lookup, ProcessEnvironment}; +use litellm_host_python::{from_py, json_object_field, run_async_value, run_sync_value, to_py}; +use litellm_secrets::{ + KeyManagementSettings, KeyManagementSystem, Secret, SecretManager, load_native_manager, + read_secret_from_python_manager, +}; +use litellm_secrets_types::PythonSecretRead; +use pyo3::{ + exceptions::{PyAttributeError, PyRuntimeError, PyValueError}, + prelude::*, +}; + +#[derive(Clone, PartialEq)] +struct Configuration { + system: KeyManagementSystem, + settings: KeyManagementSettings, + environment: BTreeMap, + enterprise_enabled: bool, +} + +#[pyclass(frozen, name = "_SecretManagerRuntime")] +pub(crate) struct NativeSecretManager { + backend: SecretManager, + configuration: Configuration, + pid: u32, +} + +impl NativeSecretManager { + pub(super) fn backend(&self) -> PyResult { + if self.pid != std::process::id() { + return Err(PyRuntimeError::new_err( + "native secret manager must be recreated after fork", + )); + } + Ok(self.backend.clone()) + } + + fn build(py: Python<'_>, configuration: Configuration) -> PyResult { + let values = configuration.environment.clone(); + let environment: Arc = + Arc::new(move |name: &str| values.get(name).cloned()); + let system = configuration.system; + let settings = configuration.settings.clone(); + let enterprise_enabled = configuration.enterprise_enabled; + let backend = run_sync_value(py, async move { + load_native_manager(system, settings, environment, enterprise_enabled) + .await + .map_err(|error| PyValueError::new_err(error.to_string())) + })?; + Ok(Self { + backend, + configuration, + pid: std::process::id(), + }) + } +} + +#[pymethods] +impl NativeSecretManager { + #[staticmethod] + #[pyo3(signature = (system, environment, settings=None, enterprise_enabled=false))] + fn from_config( + py: Python<'_>, + system: &str, + environment: BTreeMap, + settings: Option<&Bound<'_, PyAny>>, + enterprise_enabled: bool, + ) -> PyResult { + let system = serde_json::from_value(serde_json::Value::String(system.to_owned())) + .map_err(|_| PyValueError::new_err("unknown secret manager system"))?; + let settings = parse_settings(settings)?; + Self::build( + py, + Configuration { + system, + settings, + environment, + enterprise_enabled, + }, + ) + } + + #[staticmethod] + pub(super) fn from_client(client: &Bound<'_, PyAny>) -> PyResult>> { + let py = client.py(); + if let Ok(native) = client.extract::>() { + native.borrow(py).backend()?; + return Ok(Some(native)); + } + let config = py + .import("litellm.rust_bridge.secret_manager")? + .getattr("native_secret_manager_config")? + .call1((client,))?; + if config.is_none() { + return Ok(None); + } + if !config.getattr("owner_type")?.is(client.get_type()) { + return Ok(None); + } + let methods = config + .getattr("methods")? + .extract::)>>()?; + for (name, original) in methods { + let current = client.getattr(name.as_str())?; + let implementation = optional_attribute(¤t, "__func__")?.unwrap_or(current); + if !implementation.is(original.bind(py)) { + return Ok(None); + } + } + let environment_attributes: BTreeMap = config + .getattr("environment_attributes")? + .extract::>()? + .into_iter() + .collect(); + let captured = config + .getattr("environment")? + .extract::>()?; + let overrides = environment_attributes + .iter() + .map(|(key, attribute)| { + let value = attribute_path(client, attribute)?; + Ok(( + key.clone(), + if value.is_none() { + None + } else { + Some(value.str()?.extract::()?) + }, + )) + }) + .collect::>>()?; + let settings = + from_py::>(&config.getattr("settings")?) + .map_err(|_| PyValueError::new_err("invalid secret manager settings"))?; + let attributes = config + .getattr("settings_attributes")? + .extract::>()?; + let setting_overrides = attributes + .into_iter() + .map(|name| { + let value = from_py::(&client.getattr(name.as_str())?)?; + Ok((name, value)) + }) + .collect::>>()?; + let configuration = Configuration { + system: serde_json::from_value(serde_json::Value::String( + config.getattr("system")?.extract()?, + )) + .map_err(|_| PyValueError::new_err("unknown secret manager system"))?, + settings: serde_json::from_value(serde_json::Value::Object( + settings.into_iter().chain(setting_overrides).collect(), + )) + .map_err(|_| PyValueError::new_err("invalid secret manager settings"))?, + environment: captured + .into_iter() + .filter(|(key, _)| !environment_attributes.contains_key(key)) + .chain( + overrides + .into_iter() + .filter_map(|(key, value)| value.map(|value| (key, value))), + ) + .collect(), + enterprise_enabled: config.getattr("enterprise_enabled")?.extract()?, + }; + if let Some(native) = cached(client, &configuration)? { + return Ok(Some(native)); + } + let runtime = Self::build(py, configuration)?; + if let Some(native) = cached(client, &runtime.configuration)? { + return Ok(Some(native)); + } + let native = Py::new(py, runtime)?; + client.setattr("_litellm_native_secret_manager", native.bind(py))?; + Ok(Some(native)) + } + + #[getter] + fn system(&self) -> String { + serde_json::to_value(self.configuration.system) + .expect("serializable system") + .as_str() + .expect("string system") + .to_owned() + } + + #[pyo3(signature = (name, settings=None))] + fn read_secret( + &self, + py: Python<'_>, + name: String, + settings: Option<&Bound<'_, PyAny>>, + ) -> PyResult> { + let backend = self.backend()?; + let settings = settings + .map(|value| parse_settings(Some(value))) + .transpose()? + .unwrap_or_else(|| self.configuration.settings.clone()); + run_sync_value(py, async move { + read_secret_from_python_manager(&backend, &name, &settings, &ProcessEnvironment) + .await + .map_err(|error| python_read_error(backend.system(), &name, error)) + .and_then(|value| python_secret_value(value, &name)) + }) + } + #[pyo3(signature = (secret_name, optional_params=None, timeout=None, primary_secret_name=None))] + fn sync_read_secret( + &self, + py: Python<'_>, + secret_name: String, + optional_params: Option<&Bound<'_, PyAny>>, + timeout: Option<&Bound<'_, PyAny>>, + primary_secret_name: Option, + ) -> PyResult> { + let backend = self.backend()?; + let request = super::provider::read_request( + self.configuration.system, + secret_name, + optional_params, + timeout, + primary_secret_name, + true, + )?; + run_sync_value(py, async move { + super::operations::read_python_provider(&backend, &request, &ProcessEnvironment) + .await + .map_err(|error| PyValueError::new_err(error.to_string())) + .and_then(|value| python_secret_value(value, &request.secret_name)) + }) + } + + #[pyo3(signature = (secret_name, optional_params=None, timeout=None, primary_secret_name=None))] + fn async_read_secret<'py>( + &self, + py: Python<'py>, + secret_name: String, + optional_params: Option<&Bound<'py, PyAny>>, + timeout: Option<&Bound<'py, PyAny>>, + primary_secret_name: Option, + ) -> PyResult> { + let backend = self.backend()?; + let request = super::provider::read_request( + self.configuration.system, + secret_name, + optional_params, + timeout, + primary_secret_name, + false, + )?; + run_async_value(py, async move { + super::operations::read_python_provider(&backend, &request, &ProcessEnvironment) + .await + .map_err(|error| PyValueError::new_err(error.to_string())) + .and_then(|value| python_secret_value(value, &request.secret_name)) + }) + } + + #[pyo3(signature = (secret_name, secret_value, description=None, optional_params=None, timeout=None, tags=None))] + #[expect( + clippy::too_many_arguments, + reason = "preserves the Python secret-manager write signature" + )] + fn async_write_secret<'py>( + &self, + py: Python<'py>, + secret_name: String, + secret_value: String, + description: Option<&Bound<'py, PyAny>>, + optional_params: Option<&Bound<'py, PyAny>>, + timeout: Option<&Bound<'py, PyAny>>, + tags: Option<&Bound<'py, PyAny>>, + ) -> PyResult> { + let backend = self.backend()?; + let _ = tags; + let context = litellm_secrets_types::SecretWriteContext { + operation: super::provider::mutation_context( + self.configuration.system, + optional_params, + timeout, + )?, + description: if self.configuration.system == KeyManagementSystem::HashicorpVault { + description + .filter(|value| !value.is_none()) + .map(|value| { + if value.is_truthy()? { + value.extract().map(Some) + } else { + Ok(None) + } + }) + .transpose()? + .flatten() + } else { + None + }, + ..litellm_secrets_types::SecretWriteContext::default() + }; + let error_context = + super::vault::ErrorContext::capture(py, self.configuration.system, timeout)?; + run_async_value(py, async move { + super::mutation::mutation_value( + super::operations::write_python_provider_with_context( + &backend, + &secret_name, + &litellm_secrets::SecretValue::new(secret_value), + &context, + ) + .await, + &error_context, + ) + }) + } + + #[pyo3(signature = (secret_name, recovery_window_in_days=None, optional_params=None, timeout=None))] + fn async_delete_secret<'py>( + &self, + py: Python<'py>, + secret_name: String, + recovery_window_in_days: Option<&Bound<'py, PyAny>>, + optional_params: Option<&Bound<'py, PyAny>>, + timeout: Option<&Bound<'py, PyAny>>, + ) -> PyResult> { + let backend = self.backend()?; + let _ = recovery_window_in_days; + let context = + super::provider::mutation_context(self.configuration.system, optional_params, timeout)?; + let error_context = + super::vault::ErrorContext::capture(py, self.configuration.system, timeout)?; + run_async_value(py, async move { + super::mutation::mutation_value( + super::operations::delete_python_provider_with_context( + &backend, + &secret_name, + &context, + ) + .await, + &error_context, + ) + }) + } + + #[pyo3(signature = (current_secret_name, new_secret_name, new_secret_value, optional_params=None, timeout=None))] + fn async_rotate_secret<'py>( + &self, + py: Python<'py>, + current_secret_name: String, + new_secret_name: String, + new_secret_value: String, + optional_params: Option<&Bound<'py, PyAny>>, + timeout: Option<&Bound<'py, PyAny>>, + ) -> PyResult> { + let backend = self.backend()?; + let context = + super::provider::mutation_context(self.configuration.system, optional_params, timeout)?; + let error_context = + super::vault::ErrorContext::capture(py, self.configuration.system, timeout)?; + run_async_value(py, async move { + super::mutation::mutation_value( + super::operations::rotate_python_provider_with_context( + &backend, + ¤t_secret_name, + &new_secret_name, + &litellm_secrets::SecretValue::new(new_secret_value), + &context, + ) + .await, + &error_context, + ) + }) + } + + #[pyo3(signature = (name, settings=None))] + fn read_secret_async<'py>( + &self, + py: Python<'py>, + name: String, + settings: Option<&Bound<'py, PyAny>>, + ) -> PyResult> { + let backend = self.backend()?; + let settings = settings + .map(|value| parse_settings(Some(value))) + .transpose()? + .unwrap_or_else(|| self.configuration.settings.clone()); + run_async_value(py, async move { + read_secret_from_python_manager(&backend, &name, &settings, &ProcessEnvironment) + .await + .map_err(|error| python_read_error(backend.system(), &name, error)) + .and_then(|value| python_secret_value(value, &name)) + }) + } +} + +fn optional_attribute<'py>( + object: &Bound<'py, PyAny>, + name: &str, +) -> PyResult>> { + match object.getattr(name) { + Ok(value) => Ok(Some(value)), + Err(error) if error.is_instance_of::(object.py()) => Ok(None), + Err(error) => Err(error), + } +} + +fn attribute_path<'py>(object: &Bound<'py, PyAny>, path: &str) -> PyResult> { + match path.split_once('.') { + Some((head, tail)) => attribute_path(&object.getattr(head)?, tail), + None => object.getattr(path), + } +} + +fn cached( + client: &Bound<'_, PyAny>, + configuration: &Configuration, +) -> PyResult>> { + let Some(value) = optional_attribute(client, "_litellm_native_secret_manager")? else { + return Ok(None); + }; + let native = value.extract::>()?; + let same_configuration = native.borrow(client.py()).pid == std::process::id() + && &native.borrow(client.py()).configuration == configuration; + Ok(same_configuration.then_some(native)) +} + +fn parse_settings(value: Option<&Bound<'_, PyAny>>) -> PyResult { + value + .map(|value| { + serde_json::from_value(from_py::(value)?) + .map_err(|_| PyValueError::new_err("invalid secret manager settings")) + }) + .transpose() + .map(Option::unwrap_or_default) +} + +fn python_secret_value(payload: PythonSecretRead, name: &str) -> PyResult> { + let value = match payload { + PythonSecretRead::Value(value) => value, + PythonSecretRead::PrimaryJson(document) => { + return Python::attach(|py| json_object_field(py, document.expose(), name)); + } + }; + let value = match value { + None => serde_json::Value::Null, + Some(Secret::String(value)) => serde_json::Value::String(value.expose().to_owned()), + Some(Secret::Bool(value)) => serde_json::Value::Bool(value), + Some(Secret::Json(value)) => value, + }; + Python::attach(|py| to_py(py, &value)) +} + +fn python_read_error( + system: KeyManagementSystem, + name: &str, + error: litellm_secrets::Error, +) -> PyErr { + let message = match (system, error) { + (KeyManagementSystem::Cyberark, litellm_secrets::Error::ManagedSecretMissing) => { + format!("No secret found in CyberArk Secret Manager for {name}") + } + (_, error) => error.to_string(), + }; + PyValueError::new_err(message) +} diff --git a/litellm-rust/crates/python-bridge/src/secrets/vault.rs b/litellm-rust/crates/python-bridge/src/secrets/vault.rs new file mode 100644 index 00000000000..de0fa95bcef --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/secrets/vault.rs @@ -0,0 +1,182 @@ +mod operation; + +pub(super) use operation::{Failure, FailureKind, FailureStage, delete, rotate, write}; + +use litellm_secrets::hashicorp::{Error, RawOperationError}; +use pyo3::prelude::*; + +use super::mutation::{error_value, http_message, json_value}; + +pub(super) fn failure_value( + py: Python<'_>, + failure: Failure, + context: &ErrorContext, +) -> PyResult> { + let message = match *failure.kind { + FailureKind::Native(RawOperationError::Http { + method, + url, + status, + body, + }) => match failure.stage { + FailureStage::Current(name) if status == 404 => { + format!("Current secret {name} not found") + } + FailureStage::Replacement(name) if status == 404 => { + format!("Failed to verify new secret {name}") + } + FailureStage::Current(_) => format!( + "HTTP error occurred while checking current secret: {}", + response_text(py, &body)? + ), + FailureStage::Replacement(_) => format!( + "HTTP error occurred while verifying new secret: {}", + response_text(py, &body)? + ), + FailureStage::Mutation => http_message(py, &method, &url, status)?, + }, + FailureKind::ValueMismatch { expected, actual } => { + let actual = json_value(py, &actual)?; + format!( + "New secret value mismatch. Expected: {}, Got: {}", + expected.expose(), + actual.bind(py).str()? + ) + } + kind => { + let message = cause_message(py, kind, context)?; + match failure.stage { + FailureStage::Current(_) => { + format!("Error checking current secret: {message}") + } + FailureStage::Replacement(_) => { + format!("Error verifying new secret: {message}") + } + FailureStage::Mutation => message, + } + } + }; + error_value(py, message) +} + +fn cause_message(py: Python<'_>, kind: FailureKind, context: &ErrorContext) -> PyResult { + Ok(match kind { + FailureKind::Native(RawOperationError::Local(error)) => error.to_string(), + FailureKind::UnsafeName(name) => { + format!("Invalid secret_name {}", name.into_pyobject(py)?.repr()?) + } + FailureKind::Native(RawOperationError::Timeout { method, elapsed }) => { + if method == "POST" { + let elapsed = py + .import("builtins")? + .call_method1("round", (elapsed.as_secs_f64(), 3))?; + let kwargs = pyo3::types::PyDict::new(py); + kwargs.set_item( + "message", + format!( + "Connection timed out. Timeout passed={}, time taken={} seconds", + context.timeout.as_deref().unwrap_or("None"), + elapsed.str()? + ), + )?; + kwargs.set_item("model", "default-model-name")?; + kwargs.set_item("llm_provider", "litellm-httpx-handler")?; + kwargs.set_item("headers", pyo3::types::PyDict::new(py))?; + py.import("litellm")? + .getattr("Timeout")? + .call((), Some(&kwargs))? + .str()? + .extract()? + } else if context.aiohttp { + "Timeout on reading data from socket".to_owned() + } else { + String::new() + } + } + FailureKind::Native(RawOperationError::Transport(source)) => { + if let Some(error) = request_error(&source) { + if error.is_timeout() { + String::new() + } else if error.is_connect() { + "All connection attempts failed".to_owned() + } else { + "HashiCorp Vault request failed".to_owned() + } + } else { + "HashiCorp Vault request failed".to_owned() + } + } + FailureKind::MissingGet(value) => { + let value = json_value(py, &value)?; + match value.bind(py).getattr("get") { + Err(error) => error.value(py).str()?.extract()?, + Ok(_) => "HashiCorp Vault response payload is malformed".to_owned(), + } + } + FailureKind::Json(body) => match json_value(py, &body) { + Err(error) => error.value(py).str()?.extract()?, + Ok(_) => "HashiCorp Vault response payload is malformed".to_owned(), + }, + FailureKind::Native(RawOperationError::Authentication { + source, + url, + certificate, + }) => { + let message = match source { + Error::LoginStatus { status } => http_message(py, "POST", &url, status)?, + error => error.to_string(), + }; + let mechanism = if certificate { "TLS cert" } else { "AppRole" }; + format!("Could not authenticate to Vault via {mechanism}: {message}") + } + FailureKind::Native(RawOperationError::Http { + method, + url, + status, + .. + }) => http_message(py, &method, &url, status)?, + FailureKind::ValueMismatch { .. } => "New secret value mismatch".to_owned(), + }) +} + +fn request_error<'a>(error: &'a (dyn std::error::Error + 'static)) -> Option<&'a reqwest::Error> { + error + .downcast_ref::() + .or_else(|| error.source().and_then(request_error)) +} + +fn response_text(py: Python<'_>, body: &[u8]) -> PyResult { + let kwargs = pyo3::types::PyDict::new(py); + kwargs.set_item("content", pyo3::types::PyBytes::new(py, body))?; + py.import("httpx")? + .getattr("Response")? + .call((200,), Some(&kwargs))? + .getattr("text")? + .extract() +} + +#[derive(Default)] +pub(super) struct ErrorContext { + timeout: Option, + aiohttp: bool, +} + +impl ErrorContext { + pub(super) fn capture( + py: Python<'_>, + system: litellm_secrets::KeyManagementSystem, + timeout: Option<&Bound<'_, PyAny>>, + ) -> PyResult { + if system != litellm_secrets::KeyManagementSystem::HashicorpVault { + return Ok(Self::default()); + } + Ok(Self { + timeout: timeout.map(|value| value.str()?.extract()).transpose()?, + aiohttp: py + .import("litellm.llms.custom_httpx.http_handler")? + .getattr("AsyncHTTPHandler")? + .call_method0("_should_use_aiohttp_transport")? + .extract()?, + }) + } +} diff --git a/litellm-rust/crates/python-bridge/src/secrets/vault/operation.rs b/litellm-rust/crates/python-bridge/src/secrets/vault/operation.rs new file mode 100644 index 00000000000..cf56beb5efe --- /dev/null +++ b/litellm-rust/crates/python-bridge/src/secrets/vault/operation.rs @@ -0,0 +1,168 @@ +use std::collections::HashMap; + +use litellm_secrets::{ + SecretValue, + hashicorp::{Error, HashicorpVault, RawOperationError}, +}; +use litellm_secrets_types::{HashicorpOperationContext, SecretWriteContext}; +use serde_json::value::RawValue; + +#[derive(Debug)] +pub(crate) enum FailureStage { + Mutation, + Current(String), + Replacement(String), +} + +#[derive(veil::Redact)] +pub(crate) enum FailureKind { + Native(RawOperationError), + UnsafeName(#[redact] String), + Json(#[redact] Vec), + MissingGet(#[redact] Vec), + ValueMismatch { + expected: SecretValue, + #[redact] + actual: Vec, + }, +} + +#[derive(Debug)] +pub(crate) struct Failure { + pub kind: Box, + pub stage: FailureStage, +} + +impl From for Failure { + fn from(kind: FailureKind) -> Self { + Self { + kind: Box::new(kind), + stage: FailureStage::Mutation, + } + } +} + +impl Failure { + fn during(self, stage: FailureStage) -> Self { + if matches!(*self.kind, FailureKind::UnsafeName(_)) { + self + } else { + Self { stage, ..self } + } + } +} + +fn native_failure(name: &str, error: RawOperationError) -> Failure { + match error { + RawOperationError::Local(Error::InvalidSecretName(_)) => { + FailureKind::UnsafeName(name.to_owned()).into() + } + error => FailureKind::Native(error).into(), + } +} + +pub(crate) async fn write( + client: &HashicorpVault, + name: &str, + value: &SecretValue, + context: &SecretWriteContext, +) -> Result, Failure> { + client + .write_raw(name, value, context) + .await + .map_err(|error| native_failure(name, error)) +} + +pub(crate) async fn delete( + client: &HashicorpVault, + name: &str, + context: &HashicorpOperationContext, +) -> Result<(), Failure> { + client + .delete_raw(name, context) + .await + .map_err(|error| native_failure(name, error)) +} + +pub(crate) async fn rotate( + client: &HashicorpVault, + current_name: &str, + new_name: &str, + value: &SecretValue, + context: &HashicorpOperationContext, +) -> Result, Failure> { + client + .read_raw(current_name, context) + .await + .map_err(|error| { + native_failure(current_name, error) + .during(FailureStage::Current(current_name.to_owned())) + })?; + let response = write( + client, + new_name, + value, + &SecretWriteContext { + description: Some(format!("Rotated from {current_name}")), + operation: context.clone(), + ..SecretWriteContext::default() + }, + ) + .await?; + let parsed: &RawValue = + serde_json::from_slice(&response).map_err(|_| FailureKind::Json(response.clone()))?; + let status = raw_object_field(parsed, "status") + .ok() + .flatten() + .and_then(|value| serde_json::from_slice::(value.get().as_bytes()).ok()); + if status.as_deref() == Some("error") { + return Ok(response); + } + let verification = client.read_raw(new_name, context).await.map_err(|error| { + native_failure(new_name, error).during(FailureStage::Replacement(new_name.to_owned())) + })?; + let parsed: &RawValue = serde_json::from_slice(&verification).map_err(|_| { + Failure::from(FailureKind::Json(verification.clone())) + .during(FailureStage::Replacement(new_name.to_owned())) + })?; + let data_key = context + .data_key + .as_deref() + .map(str::trim) + .filter(|key| !key.is_empty()) + .unwrap_or("key"); + let actual = verification_value(parsed, data_key).map_err(|failure| { + Failure::from(failure).during(FailureStage::Replacement(new_name.to_owned())) + })?; + let actual_string = serde_json::from_slice::(actual.get().as_bytes()).ok(); + if actual_string.as_deref() != Some(value.expose()) { + return Err(FailureKind::ValueMismatch { + expected: value.clone(), + actual: actual.get().as_bytes().to_vec(), + } + .into()); + } + if current_name != new_name { + let _ = delete(client, current_name, context).await; + } + Ok(response) +} + +fn verification_value<'a>(document: &'a RawValue, key: &str) -> Result<&'a RawValue, FailureKind> { + let Some(outer) = raw_object_field(document, "data")? else { + return Ok(RawValue::NULL); + }; + let Some(inner) = raw_object_field(outer, "data")? else { + return Ok(RawValue::NULL); + }; + Ok(raw_object_field(inner, key)?.unwrap_or(RawValue::NULL)) +} + +fn raw_object_field<'a>( + document: &'a RawValue, + key: &str, +) -> Result, FailureKind> { + let object: HashMap = serde_json::from_slice(document.get().as_bytes()) + .map_err(|_| FailureKind::MissingGet(document.get().as_bytes().to_vec()))?; + Ok(object.get(key).copied()) +} diff --git a/litellm-rust/crates/secrets-aws/AGENTS.md b/litellm-rust/crates/secrets-aws/AGENTS.md new file mode 100644 index 00000000000..714f031db41 --- /dev/null +++ b/litellm-rust/crates/secrets-aws/AGENTS.md @@ -0,0 +1 @@ +- https://docs.aws.amazon.com/secretsmanager/latest/apireference/Welcome.html diff --git a/litellm-rust/crates/secrets-aws/Cargo.toml b/litellm-rust/crates/secrets-aws/Cargo.toml index 9c926bb441b..e7a394bd247 100644 --- a/litellm-rust/crates/secrets-aws/Cargo.toml +++ b/litellm-rust/crates/secrets-aws/Cargo.toml @@ -22,3 +22,4 @@ base64.workspace = true rstest.workspace = true tokio.workspace = true wiremock = "0.6.5" +tempfile = "3" diff --git a/litellm-rust/crates/secrets-aws/src/auth.rs b/litellm-rust/crates/secrets-aws/src/auth.rs index 954cfa2f8fd..0c32eb00989 100644 --- a/litellm-rust/crates/secrets-aws/src/auth.rs +++ b/litellm-rust/crates/secrets-aws/src/auth.rs @@ -7,7 +7,7 @@ use litellm_auth_aws::{ resolve_credentials, }; use litellm_core_utils::settings::Lookup; -use litellm_secrets_types::KeyManagementSettings; +use litellm_secrets_types::{AwsOperationContext, KeyManagementSettings}; use crate::Error; @@ -21,9 +21,29 @@ impl Credentials { pub(crate) fn new( settings: &KeyManagementSettings, environment: Arc, + ) -> Self { + Self::with_context(settings, environment, &AwsOperationContext::default()) + } + + pub(crate) fn with_context( + settings: &KeyManagementSettings, + environment: Arc, + context: &AwsOperationContext, ) -> Self { Self { config: AwsAuthConfig { + access_key_id: context + .access_key_id + .as_ref() + .map(|value| value.expose().to_owned()), + secret_access_key: context + .secret_access_key + .as_ref() + .map(|value| value.expose().to_owned()), + session_token: context + .session_token + .as_ref() + .map(|value| value.expose().to_owned()), region_name: region(settings, environment.as_ref()).ok(), role_name: settings.aws_role_name.clone(), session_name: settings.aws_session_name.clone(), @@ -37,7 +57,6 @@ impl Credentials { .as_ref() .map(|v| v.expose().to_owned()), sts_endpoint: settings.aws_sts_endpoint.clone(), - ..Default::default() }, environment, } diff --git a/litellm-rust/crates/secrets-aws/src/error.rs b/litellm-rust/crates/secrets-aws/src/error.rs index cc9e6b69786..c8ca671e9b5 100644 --- a/litellm-rust/crates/secrets-aws/src/error.rs +++ b/litellm-rust/crates/secrets-aws/src/error.rs @@ -6,8 +6,6 @@ pub enum Error { Auth(#[from] #[redact] litellm_auth_aws::Error), #[error("AWS region is not configured")] MissingRegion, - #[error("AWS Secrets Manager received a non-AWS operation context")] - InvalidOperationContext, #[error("AWS Secrets Manager was constructed without context-aware configuration")] OperationContextUnavailable, #[error("KMS response has no plaintext")] @@ -20,6 +18,12 @@ pub enum Error { Read(#[from] #[redact] Box>), #[error("AWS Secrets Manager create failed")] Create(#[from] #[redact] Box>), + #[error("AWS Secrets Manager restore failed")] + Restore(#[from] #[redact] Box>), + #[error("AWS Secrets Manager restored update failed")] + Update(#[from] #[redact] Box>), + #[error("AWS Secrets Manager tagging failed")] + Tag(#[from] #[redact] Box>), #[error("AWS Secrets Manager update failed")] Put(#[from] #[redact] Box>), #[error("AWS Secrets Manager delete failed")] diff --git a/litellm-rust/crates/secrets-aws/src/kms.rs b/litellm-rust/crates/secrets-aws/src/kms.rs index a66b1c4d2fe..8c05c016184 100644 --- a/litellm-rust/crates/secrets-aws/src/kms.rs +++ b/litellm-rust/crates/secrets-aws/src/kms.rs @@ -1,4 +1,3 @@ -use litellm_auth_aws::constants::AWS_REGION_NAME; use std::sync::Arc; use aws_sdk_kms::{ @@ -37,10 +36,7 @@ impl AwsKms { } pub fn validate_environment(environment: &dyn Lookup) -> Result<(), Error> { - environment - .get(AWS_REGION_NAME) - .map(|_| ()) - .ok_or(Error::MissingRegion) + auth::region(&KeyManagementSettings::default(), environment).map(|_| ()) } pub fn load_aws_kms( @@ -51,9 +47,6 @@ pub fn load_aws_kms( if use_aws_kms != Some(true) { return Ok(None); } - if settings.aws_region_name.is_none() { - validate_environment(environment.as_ref())?; - } let config = aws_sdk_kms::Config::builder() .behavior_version(BehaviorVersion::latest()) .region(Region::new(auth::region(settings, environment.as_ref())?)) diff --git a/litellm-rust/crates/secrets-aws/src/secret_manager.rs b/litellm-rust/crates/secrets-aws/src/secret_manager.rs index 5658b6d7e11..508d7da15c1 100644 --- a/litellm-rust/crates/secrets-aws/src/secret_manager.rs +++ b/litellm-rust/crates/secrets-aws/src/secret_manager.rs @@ -1,3 +1,9 @@ +mod client; +mod read; +mod write; + +pub use read::is_bootstrap_key; + use litellm_auth_aws::constants::AWS_BEDROCK_RUNTIME_ENDPOINT; use std::{collections::BTreeMap, sync::Arc}; @@ -16,8 +22,9 @@ use litellm_auth_aws::constants::{ }; use litellm_core_utils::settings::Lookup; use litellm_secrets_types::{ - AwsOperationContext, BaseSecretManager, KeyManagementSettings, Secret, SecretOperationContext, - SecretValue, SecretWriteContext, async_rotate_secret, + AwsOperationContext, BaseSecretManager, KeyManagementSettings, RotationError, Secret, + SecretDeleter, SecretRotator, SecretValue, SecretWriteContext, SecretWriter, + async_rotate_secret, }; use serde_json::Value; @@ -68,395 +75,4 @@ impl AwsSecretsManagerV2 { write_settings, } } - - fn with_context_client_factory( - client: Client, - write_settings: AwsSecretWriteSettings, - context_client_factory: ContextClientFactory, - ) -> Self { - Self { - client, - context_client_factory: Some(Box::new(context_client_factory)), - write_settings, - } - } - - pub fn load_aws_secret_manager( - use_aws_secret_manager: Option, - settings: KeyManagementSettings, - environment: Arc, - ) -> Result, Error> { - if use_aws_secret_manager != Some(true) { - return Ok(None); - } - let context_client_factory = ContextClientFactory { - settings: settings.clone(), - environment: environment.clone(), - endpoint_url: environment - .get(AWS_BEDROCK_RUNTIME_ENDPOINT) - .map(|url| url.replace("bedrock-runtime", "secretsmanager")), - }; - let client = context_client_factory.client(&AwsOperationContext::default())?; - Ok(Some(Self::with_context_client_factory( - client, - (&settings).into(), - context_client_factory, - ))) - } - - pub async fn read_secret_for_resolver( - &self, - name: &str, - primary_name: Option<&str>, - environment: &(dyn Lookup + Sync), - ) -> Result, Error> { - if bootstrap_key(name) { - return Ok(environment - .get(name) - .map(SecretValue::new) - .map(Secret::String)); - } - match primary_name.filter(|name| !name.is_empty()) { - None => self - .async_read_secret(name) - .await - .map(|value| value.map(Secret::String)), - Some(primary) => { - let value = if bootstrap_key(primary) { - environment.get(primary).map(SecretValue::new) - } else { - self.async_read_secret(primary).await? - }; - let Some(value) = value else { - return Ok(None); - }; - let object: Value = - serde_json::from_str(value.expose()).map_err(|_| Error::PrimarySecret)?; - let object = object.as_object().ok_or(Error::PrimarySecret)?; - Ok(object.get(name).cloned().map(Secret::from_json)) - } - } - } - - pub async fn async_read_secret(&self, name: &str) -> Result, Error> { - Self::async_read_secret_with_client(&self.client, name).await - } - - async fn async_read_secret_with_client( - client: &Client, - name: &str, - ) -> Result, Error> { - match client.get_secret_value().secret_id(name).send().await { - Ok(response) => response - .secret_string - .map(SecretValue::new) - .map(Some) - .ok_or(Error::MissingString), - Err(error) - if matches!( - &error, - aws_sdk_secretsmanager::error::SdkError::TimeoutError(_) - ) || matches!(&error, aws_sdk_secretsmanager::error::SdkError::DispatchFailure(failure) if failure.is_timeout()) => - { - Err(Error::Timeout) - } - Err(error) - if error - .as_service_error() - .is_some_and(|error| error.is_resource_not_found_exception()) => - { - Ok(None) - } - Err(error) => Err(Error::Read(Box::new(error))), - } - } - - pub async fn async_write_secret( - &self, - name: &str, - value: &SecretValue, - description: Option<&str>, - ) -> Result { - self.async_write_secret_with_client_and_tags(&self.client, name, value, description, None) - .await - } - - async fn async_write_secret_with_client_and_tags( - &self, - client: &Client, - name: &str, - value: &SecretValue, - description: Option<&str>, - tags: Option<&BTreeMap>, - ) -> Result { - let response = client - .create_secret() - .name(name) - .secret_string(value.expose()) - .set_description(description.filter(|v| !v.is_empty()).map(str::to_owned)) - .set_kms_key_id( - self.write_settings - .kms_key_id - .clone() - .filter(|v| !v.is_empty()), - ) - .set_tags(tags.or(self.write_settings.tags.as_ref()).map(|tags| { - tags.iter() - .map(|(key, value)| Tag::builder().key(key).value(value).build()) - .collect() - })) - .send() - .await - .map_err(|error| Error::Create(Box::new(error)))?; - if let Some(regions) = &self.write_settings.replica_regions - && !regions.is_empty() - && self - .async_replicate_secret_with_client(client, name, regions) - .await - .is_err() - { - litellm_tracing::warn!("secret created but replication failed"); - } - Ok(response) - } - - pub async fn async_replicate_secret( - &self, - name: &str, - regions: &[String], - ) -> Result, Error> { - self.async_replicate_secret_with_client(&self.client, name, regions) - .await - } - - async fn async_replicate_secret_with_client( - &self, - client: &Client, - name: &str, - regions: &[String], - ) -> Result, Error> { - if regions.is_empty() { - return Ok(None); - } - client - .replicate_secret_to_regions() - .secret_id(name) - .set_add_replica_regions(Some( - regions - .iter() - .map(|region| ReplicaRegionType::builder().region(region).build()) - .collect(), - )) - .send() - .await - .map(Some) - .map_err(|error| Error::Replicate(Box::new(error))) - } - - pub async fn async_put_secret_value( - &self, - name: &str, - value: &SecretValue, - ) -> Result { - self.async_put_secret_value_with_client(&self.client, name, value) - .await - } - - async fn async_put_secret_value_with_client( - &self, - client: &Client, - name: &str, - value: &SecretValue, - ) -> Result { - client - .put_secret_value() - .secret_id(name) - .secret_string(value.expose()) - .send() - .await - .map_err(|error| Error::Put(Box::new(error))) - } - - pub async fn async_delete_secret( - &self, - name: &str, - recovery_window_in_days: Option, - ) -> Result { - self.async_delete_secret_with_client(&self.client, name, recovery_window_in_days) - .await - } - - async fn async_delete_secret_with_client( - &self, - client: &Client, - name: &str, - recovery_window_in_days: Option, - ) -> Result { - client - .delete_secret() - .secret_id(name) - .set_recovery_window_in_days(recovery_window_in_days.map(i64::from)) - .send() - .await - .map_err(|error| Error::Delete(Box::new(error))) - } - - pub async fn async_rotate_secret( - &self, - current_name: &str, - new_name: &str, - value: &SecretValue, - ) -> Result { - self.async_rotate_secret_with_context( - current_name, - new_name, - value, - &SecretOperationContext::default(), - ) - .await - } - - pub async fn async_rotate_secret_with_context( - &self, - current_name: &str, - new_name: &str, - value: &SecretValue, - context: &SecretOperationContext, - ) -> Result { - if current_name == new_name { - let client = self.client_for_context(context)?; - return self - .async_put_secret_value_with_client(&client, current_name, value) - .await - .map(RotationResponse::Updated); - } - async_rotate_secret(self, current_name, new_name, value, context) - .await - .map(RotationResponse::Created) - } - - fn client_for_context(&self, context: &SecretOperationContext) -> Result { - match context { - SecretOperationContext::Default => Ok(self.client.clone()), - SecretOperationContext::Aws(context) if context == &AwsOperationContext::default() => { - Ok(self.client.clone()) - } - SecretOperationContext::Aws(context) => self - .context_client_factory - .as_ref() - .ok_or(Error::OperationContextUnavailable)? - .client(context), - _ => Err(Error::InvalidOperationContext), - } - } -} - -impl ContextClientFactory { - fn client(&self, context: &AwsOperationContext) -> Result { - let settings = KeyManagementSettings { - aws_region_name: context - .region_name - .clone() - .or_else(|| self.settings.aws_region_name.clone()), - aws_role_name: context - .role_name - .clone() - .or_else(|| self.settings.aws_role_name.clone()), - aws_session_name: context - .session_name - .clone() - .or_else(|| self.settings.aws_session_name.clone()), - aws_external_id: context - .external_id - .clone() - .or_else(|| self.settings.aws_external_id.clone()), - aws_profile_name: context - .profile_name - .clone() - .or_else(|| self.settings.aws_profile_name.clone()), - aws_web_identity_token: context - .web_identity_token - .clone() - .or_else(|| self.settings.aws_web_identity_token.clone()), - aws_sts_endpoint: context - .sts_endpoint - .clone() - .or_else(|| self.settings.aws_sts_endpoint.clone()), - ..self.settings.clone() - }; - let builder = aws_sdk_secretsmanager::Config::builder() - .behavior_version(BehaviorVersion::latest()) - .region(Region::new(auth::region( - &settings, - self.environment.as_ref(), - )?)) - .credentials_provider(auth::Credentials::new(&settings, self.environment.clone())); - let builder = match context.timeout { - Some(timeout) => builder.timeout_config( - aws_sdk_secretsmanager::config::timeout::TimeoutConfig::builder() - .operation_timeout(timeout) - .build(), - ), - None => builder, - }; - let config = match &self.endpoint_url { - Some(endpoint_url) => builder.endpoint_url(endpoint_url.clone()).build(), - None => builder.build(), - }; - Ok(Client::from_conf(config)) - } -} - -impl BaseSecretManager for AwsSecretsManagerV2 { - type Error = Error; - type WriteResponse = CreateSecretOutput; - type DeleteResponse = DeleteSecretOutput; - - async fn async_read_secret( - &self, - name: &str, - context: &SecretOperationContext, - ) -> Result, Error> { - let client = self.client_for_context(context)?; - Self::async_read_secret_with_client(&client, name).await - } - - async fn async_write_secret( - &self, - name: &str, - value: &SecretValue, - context: &SecretWriteContext, - ) -> Result { - let client = self.client_for_context(&context.operation)?; - self.async_write_secret_with_client_and_tags( - &client, - name, - value, - context.description.as_deref(), - (!context.tags.is_empty()).then_some(&context.tags), - ) - .await - } - - async fn async_delete_secret( - &self, - name: &str, - recovery_window_in_days: Option, - context: &SecretOperationContext, - ) -> Result { - let client = self.client_for_context(context)?; - self.async_delete_secret_with_client(&client, name, recovery_window_in_days) - .await - } -} - -fn bootstrap_key(name: &str) -> bool { - matches!( - name, - AWS_ACCESS_KEY_ID - | AWS_SECRET_ACCESS_KEY - | AWS_REGION_NAME - | AWS_REGION - | AWS_BEDROCK_RUNTIME_ENDPOINT - ) } diff --git a/litellm-rust/crates/secrets-aws/src/secret_manager/client.rs b/litellm-rust/crates/secrets-aws/src/secret_manager/client.rs new file mode 100644 index 00000000000..aac998c65ab --- /dev/null +++ b/litellm-rust/crates/secrets-aws/src/secret_manager/client.rs @@ -0,0 +1,116 @@ +use super::*; + +impl AwsSecretsManagerV2 { + pub(super) fn with_context_client_factory( + client: Client, + write_settings: AwsSecretWriteSettings, + context_client_factory: ContextClientFactory, + ) -> Self { + Self { + client, + context_client_factory: Some(Box::new(context_client_factory)), + write_settings, + } + } + + pub fn load_aws_secret_manager( + use_aws_secret_manager: Option, + settings: KeyManagementSettings, + environment: Arc, + ) -> Result, Error> { + if use_aws_secret_manager != Some(true) { + return Ok(None); + } + let context_client_factory = ContextClientFactory { + settings: settings.clone(), + environment: environment.clone(), + endpoint_url: environment + .get(AWS_BEDROCK_RUNTIME_ENDPOINT) + .map(|url| url.replace("bedrock-runtime", "secretsmanager")), + }; + let client = context_client_factory.client(&AwsOperationContext::default())?; + Ok(Some(Self::with_context_client_factory( + client, + (&settings).into(), + context_client_factory, + ))) + } + + pub(super) fn client_for_context( + &self, + context: &AwsOperationContext, + ) -> Result { + if context == &AwsOperationContext::default() { + return Ok(self.client.clone()); + } + self.context_client_factory + .as_ref() + .ok_or(Error::OperationContextUnavailable)? + .client(context) + } +} + +impl ContextClientFactory { + fn client(&self, context: &AwsOperationContext) -> Result { + let settings = KeyManagementSettings { + aws_region_name: context + .region_name + .clone() + .or_else(|| self.settings.aws_region_name.clone()), + aws_role_name: context + .role_name + .clone() + .or_else(|| self.settings.aws_role_name.clone()), + aws_session_name: context + .session_name + .clone() + .or_else(|| self.settings.aws_session_name.clone()), + aws_external_id: context + .external_id + .clone() + .or_else(|| self.settings.aws_external_id.clone()), + aws_profile_name: context + .profile_name + .clone() + .or_else(|| self.settings.aws_profile_name.clone()), + aws_web_identity_token: context + .web_identity_token + .clone() + .or_else(|| self.settings.aws_web_identity_token.clone()), + aws_sts_endpoint: context + .sts_endpoint + .clone() + .or_else(|| self.settings.aws_sts_endpoint.clone()), + ..self.settings.clone() + }; + let builder = aws_sdk_secretsmanager::Config::builder() + .behavior_version(BehaviorVersion::latest()) + .region(Region::new(auth::region( + &settings, + self.environment.as_ref(), + )?)) + .credentials_provider(auth::Credentials::with_context( + &settings, + self.environment.clone(), + context, + )); + let builder = match context.timeout { + Some(timeout) => builder.timeout_config( + aws_sdk_secretsmanager::config::timeout::TimeoutConfig::builder() + .operation_timeout(timeout) + .build(), + ), + None => builder, + }; + let endpoint_url = context + .bedrock_runtime_endpoint + .as_ref() + .map(|url| url.replace("bedrock-runtime", "secretsmanager")) + .or_else(|| self.endpoint_url.clone()); + let config = match endpoint_url { + Some(endpoint_url) => builder.endpoint_url(endpoint_url).build(), + None => builder.build(), + }; + Ok(Client::from_conf(config)) + } +} diff --git a/litellm-rust/crates/secrets-aws/src/secret_manager/read.rs b/litellm-rust/crates/secrets-aws/src/secret_manager/read.rs new file mode 100644 index 00000000000..b7583f9785f --- /dev/null +++ b/litellm-rust/crates/secrets-aws/src/secret_manager/read.rs @@ -0,0 +1,225 @@ +use super::*; +use aws_sdk_secretsmanager::config::retry::RetryConfig; +use litellm_secrets_types::PythonSecretRead; + +#[derive(Clone, Copy)] +enum ReadPolicy { + Native, + Python, +} + +impl AwsSecretsManagerV2 { + pub async fn read_secret_for_resolver( + &self, + name: &str, + primary_name: Option<&str>, + environment: &(dyn Lookup + Sync), + ) -> Result, Error> { + let payload = self + .read_payload(name, primary_name, environment, ReadPolicy::Native) + .await?; + resolve_payload(payload, name) + } + + pub async fn read_secret_for_python( + &self, + name: &str, + primary_name: Option<&str>, + environment: &(dyn Lookup + Sync), + ) -> Result, Error> { + let payload = self + .read_payload_for_python(name, primary_name, environment) + .await?; + resolve_payload(payload, name) + } + + pub async fn read_payload_for_python( + &self, + name: &str, + primary_name: Option<&str>, + environment: &(dyn Lookup + Sync), + ) -> Result { + self.read_payload(name, primary_name, environment, ReadPolicy::Python) + .await + } + + pub async fn read_provider_payload_for_python( + &self, + name: &str, + primary_name: Option<&str>, + context: &AwsOperationContext, + synchronous: bool, + environment: &(dyn Lookup + Sync), + ) -> Result { + if synchronous && is_bootstrap_key(name) { + return Ok(PythonSecretRead::Value( + environment + .get(name) + .map(SecretValue::new) + .map(Secret::String), + )); + } + if let Some(primary) = primary_name.filter(|value| !value.is_empty()) { + let value = if synchronous && is_bootstrap_key(primary) { + environment.get(primary).map(SecretValue::new) + } else { + self.read_with_policy(primary, ReadPolicy::Python).await? + }; + return Ok(match value.filter(|value| !value.expose().is_empty()) { + Some(value) => PythonSecretRead::PrimaryJson(value), + None => PythonSecretRead::Value(None), + }); + } + let client = self.client_for_context(context)?; + let value = match Self::read_with_client(&client, name, ReadPolicy::Python).await { + Err(Error::Read(_) | Error::MissingString | Error::Timeout) => None, + result => result?, + }; + Ok(PythonSecretRead::Value(value.map(Secret::String))) + } + + async fn read_payload( + &self, + name: &str, + primary_name: Option<&str>, + environment: &(dyn Lookup + Sync), + policy: ReadPolicy, + ) -> Result { + if is_bootstrap_key(name) { + return Ok(PythonSecretRead::Value( + environment + .get(name) + .map(SecretValue::new) + .map(Secret::String), + )); + } + match primary_name.filter(|name| !name.is_empty()) { + None => self + .read_with_policy(name, policy) + .await + .map(|value| PythonSecretRead::Value(value.map(Secret::String))), + Some(primary) => { + let value = if is_bootstrap_key(primary) { + environment.get(primary).map(SecretValue::new) + } else { + self.read_with_policy(primary, policy).await? + }; + let Some(value) = value else { + return Ok(PythonSecretRead::Value(None)); + }; + if matches!(policy, ReadPolicy::Python) && value.expose().is_empty() { + return Ok(PythonSecretRead::Value(None)); + } + Ok(PythonSecretRead::PrimaryJson(value)) + } + } + } + + async fn read_with_policy( + &self, + name: &str, + policy: ReadPolicy, + ) -> Result, Error> { + match ( + Self::read_with_client(&self.client, name, policy).await, + policy, + ) { + (Err(Error::Read(_) | Error::MissingString | Error::Timeout), ReadPolicy::Python) => { + Ok(None) + } + (result, _) => result, + } + } + + pub async fn async_read_secret(&self, name: &str) -> Result, Error> { + Self::async_read_secret_with_client(&self.client, name).await + } + + pub(super) async fn async_read_secret_with_client( + client: &Client, + name: &str, + ) -> Result, Error> { + Self::read_with_client(client, name, ReadPolicy::Native).await + } + + async fn read_with_client( + client: &Client, + name: &str, + policy: ReadPolicy, + ) -> Result, Error> { + let request = client.get_secret_value().secret_id(name); + let response = match policy { + ReadPolicy::Native => request.send().await, + ReadPolicy::Python => { + request + .customize() + .config_override( + aws_sdk_secretsmanager::config::Builder::new() + .retry_config(RetryConfig::disabled()), + ) + .send() + .await + } + }; + match response { + Ok(response) => response + .secret_string + .map(SecretValue::new) + .map(Some) + .ok_or(Error::MissingString), + Err(error) + if matches!( + &error, + aws_sdk_secretsmanager::error::SdkError::TimeoutError(_) + ) || matches!(&error, aws_sdk_secretsmanager::error::SdkError::DispatchFailure(failure) if failure.is_timeout()) => + { + Err(Error::Timeout) + } + Err(error) + if error + .as_service_error() + .is_some_and(|error| error.is_resource_not_found_exception()) => + { + Ok(None) + } + Err(error) => Err(Error::Read(Box::new(error))), + } + } +} + +impl BaseSecretManager for AwsSecretsManagerV2 { + type Error = Error; + type Context = AwsOperationContext; + + async fn async_read_secret( + &self, + name: &str, + context: &Self::Context, + ) -> Result, Error> { + let client = self.client_for_context(context)?; + Self::async_read_secret_with_client(&client, name).await + } +} + +pub fn is_bootstrap_key(name: &str) -> bool { + matches!( + name, + AWS_ACCESS_KEY_ID + | AWS_SECRET_ACCESS_KEY + | AWS_REGION_NAME + | AWS_REGION + | AWS_BEDROCK_RUNTIME_ENDPOINT + ) +} + +fn resolve_payload(payload: PythonSecretRead, name: &str) -> Result, Error> { + match payload { + PythonSecretRead::Value(value) => Ok(value), + PythonSecretRead::PrimaryJson(document) => { + let object: Value = + serde_json::from_str(document.expose()).map_err(|_| Error::PrimarySecret)?; + let object = object.as_object().ok_or(Error::PrimarySecret)?; + Ok(object.get(name).cloned().map(Secret::from_json)) + } + } +} diff --git a/litellm-rust/crates/secrets-aws/src/secret_manager/write.rs b/litellm-rust/crates/secrets-aws/src/secret_manager/write.rs new file mode 100644 index 00000000000..003b559d56c --- /dev/null +++ b/litellm-rust/crates/secrets-aws/src/secret_manager/write.rs @@ -0,0 +1,329 @@ +use super::*; + +impl AwsSecretsManagerV2 { + pub async fn async_write_secret( + &self, + name: &str, + value: &SecretValue, + description: Option<&str>, + ) -> Result { + self.async_write_secret_with_client_and_tags(&self.client, name, value, description, None) + .await + } + + pub(super) async fn async_write_secret_with_client_and_tags( + &self, + client: &Client, + name: &str, + value: &SecretValue, + description: Option<&str>, + tags: Option<&BTreeMap>, + ) -> Result { + let tags = self.write_tags(tags); + let request = client + .create_secret() + .name(name) + .secret_string(value.expose()) + .set_description(description.filter(|v| !v.is_empty()).map(str::to_owned)) + .set_kms_key_id(self.write_kms_key_id()) + .set_tags(tags.clone()); + let response = match request.send().await { + Ok(response) => response, + Err(error) => self + .restore_and_update_secret(client, name, value, description, tags) + .await? + .ok_or_else(|| Error::Create(Box::new(error)))?, + }; + if let Some(regions) = &self.write_settings.replica_regions + && !regions.is_empty() + && self + .async_replicate_secret_with_client(client, name, regions) + .await + .is_err() + { + litellm_tracing::warn!("secret created but replication failed"); + } + Ok(response) + } + + async fn restore_and_update_secret( + &self, + client: &Client, + name: &str, + value: &SecretValue, + description: Option<&str>, + tags: Option>, + ) -> Result, Error> { + let scheduled = client + .describe_secret() + .secret_id(name) + .send() + .await + .is_ok_and(|response| response.deleted_date().is_some()); + if !scheduled { + return Ok(None); + } + client + .restore_secret() + .secret_id(name) + .send() + .await + .map_err(|error| Error::Restore(Box::new(error)))?; + match self + .update_restored_secret(client, name, value, description, tags) + .await + { + Ok(response) => Ok(Some(response)), + Err(error) => { + self.async_delete_secret_with_client(client, name, Some(7)) + .await?; + Err(error) + } + } + } + + fn write_kms_key_id(&self) -> Option { + self.write_settings + .kms_key_id + .clone() + .filter(|value| !value.is_empty()) + } + + fn write_tags(&self, tags: Option<&BTreeMap>) -> Option> { + tags.or(self.write_settings.tags.as_ref()).map(|tags| { + tags.iter() + .map(|(key, value)| Tag::builder().key(key).value(value).build()) + .collect() + }) + } + + async fn update_restored_secret( + &self, + client: &Client, + name: &str, + value: &SecretValue, + description: Option<&str>, + tags: Option>, + ) -> Result { + let response = client + .update_secret() + .secret_id(name) + .secret_string(value.expose()) + .set_description( + description + .filter(|value| !value.is_empty()) + .map(str::to_owned), + ) + .set_kms_key_id(self.write_kms_key_id()) + .send() + .await + .map_err(|error| Error::Update(Box::new(error)))?; + if let Some(tags) = tags { + client + .tag_resource() + .secret_id(name) + .set_tags(Some(tags)) + .send() + .await + .map_err(|error| Error::Tag(Box::new(error)))?; + } + Ok(CreateSecretOutput::builder() + .set_arn(response.arn) + .set_name(response.name) + .set_version_id(response.version_id) + .build()) + } + + pub async fn async_replicate_secret( + &self, + name: &str, + regions: &[String], + ) -> Result, Error> { + self.async_replicate_secret_with_client(&self.client, name, regions) + .await + } + + pub(super) async fn async_replicate_secret_with_client( + &self, + client: &Client, + name: &str, + regions: &[String], + ) -> Result, Error> { + if regions.is_empty() { + return Ok(None); + } + client + .replicate_secret_to_regions() + .secret_id(name) + .set_add_replica_regions(Some( + regions + .iter() + .map(|region| ReplicaRegionType::builder().region(region).build()) + .collect(), + )) + .send() + .await + .map(Some) + .map_err(|error| Error::Replicate(Box::new(error))) + } + + pub async fn async_put_secret_value( + &self, + name: &str, + value: &SecretValue, + ) -> Result { + self.async_put_secret_value_with_client(&self.client, name, value) + .await + } + + pub(super) async fn async_put_secret_value_with_client( + &self, + client: &Client, + name: &str, + value: &SecretValue, + ) -> Result { + client + .put_secret_value() + .secret_id(name) + .secret_string(value.expose()) + .send() + .await + .map_err(|error| Error::Put(Box::new(error))) + } + + pub async fn async_delete_secret( + &self, + name: &str, + recovery_window_in_days: Option, + ) -> Result { + self.async_delete_secret_with_client(&self.client, name, recovery_window_in_days) + .await + } + + pub async fn async_delete_secret_with_context( + &self, + name: &str, + recovery_window_in_days: Option, + context: &AwsOperationContext, + ) -> Result { + let client = self.client_for_context(context)?; + self.async_delete_secret_with_client(&client, name, recovery_window_in_days) + .await + } + + pub(super) async fn async_delete_secret_with_client( + &self, + client: &Client, + name: &str, + recovery_window_in_days: Option, + ) -> Result { + client + .delete_secret() + .secret_id(name) + .set_recovery_window_in_days(recovery_window_in_days.map(i64::from)) + .send() + .await + .map_err(|error| Error::Delete(Box::new(error))) + } + + pub async fn async_rotate_secret( + &self, + current_name: &str, + new_name: &str, + value: &SecretValue, + ) -> Result> { + self.async_rotate_secret_with_context( + current_name, + new_name, + value, + &AwsOperationContext::default(), + ) + .await + } + + pub async fn async_rotate_secret_with_context( + &self, + current_name: &str, + new_name: &str, + value: &SecretValue, + context: &AwsOperationContext, + ) -> Result> { + if current_name == new_name { + return self + .async_write_replacement(current_name, new_name, value, context) + .await + .map_err(RotationError::Write); + } + async_rotate_secret(self, current_name, new_name, value, context).await + } +} + +impl SecretWriter for AwsSecretsManagerV2 { + type WriteResponse = CreateSecretOutput; + + async fn async_write_secret( + &self, + name: &str, + value: &SecretValue, + context: &SecretWriteContext, + ) -> Result { + let client = self.client_for_context(&context.operation)?; + self.async_write_secret_with_client_and_tags( + &client, + name, + value, + context.description.as_deref(), + (!context.tags.is_empty()).then_some(&context.tags), + ) + .await + } +} + +impl SecretDeleter for AwsSecretsManagerV2 { + type DeleteResponse = DeleteSecretOutput; + + async fn async_delete_secret( + &self, + name: &str, + context: &Self::Context, + ) -> Result { + self.async_delete_secret_with_context(name, Some(7), context) + .await + } +} + +impl SecretRotator for AwsSecretsManagerV2 { + type RotationResponse = RotationResponse; + + async fn async_read_secret_fresh( + &self, + name: &str, + context: &Self::Context, + ) -> Result, Error> { + BaseSecretManager::async_read_secret(self, name, context).await + } + + async fn async_write_replacement( + &self, + current_name: &str, + new_name: &str, + value: &SecretValue, + context: &Self::Context, + ) -> Result { + if current_name == new_name { + let client = self.client_for_context(context)?; + return self + .async_put_secret_value_with_client(&client, new_name, value) + .await + .map(RotationResponse::Updated); + } + SecretWriter::async_write_secret( + self, + new_name, + value, + &SecretWriteContext::rotated_from(current_name, context.clone()), + ) + .await + .map(RotationResponse::Created) + } +} diff --git a/litellm-rust/crates/secrets-aws/tests/kms.rs b/litellm-rust/crates/secrets-aws/tests/kms.rs index 687cca104e0..71be88e7313 100644 --- a/litellm-rust/crates/secrets-aws/tests/kms.rs +++ b/litellm-rust/crates/secrets-aws/tests/kms.rs @@ -60,10 +60,13 @@ fn disabled_kms_loader_does_not_require_environment_configuration(#[case] enable } #[rstest] -#[case::settings(Some("configured-region"), None)] -#[case::environment(None, Some("environment-region"))] -fn enabled_kms_loader_accepts_either_region_source( +#[case::settings(Some("configured-region"), None, None)] +#[case::region_name(None, Some("AWS_REGION_NAME"), Some("environment-region"))] +#[case::region(None, Some("AWS_REGION"), Some("environment-region"))] +#[case::default_region(None, Some("AWS_DEFAULT_REGION"), Some("environment-region"))] +fn enabled_kms_loader_accepts_supported_region_sources( #[case] configured_region: Option<&'static str>, + #[case] environment_region_name: Option<&'static str>, #[case] environment_region: Option<&'static str>, ) { use std::sync::Arc; @@ -72,7 +75,7 @@ fn enabled_kms_loader_accepts_either_region_source( ..KeyManagementSettings::default() }; let environment = Arc::new(move |name: &str| { - (name == "AWS_REGION_NAME") + (Some(name) == environment_region_name) .then(|| environment_region.map(str::to_owned)) .flatten() }); diff --git a/litellm-rust/crates/secrets-aws/tests/secret_manager.rs b/litellm-rust/crates/secrets-aws/tests/secret_manager.rs index 7dbc8b61374..c482a168090 100644 --- a/litellm-rust/crates/secrets-aws/tests/secret_manager.rs +++ b/litellm-rust/crates/secrets-aws/tests/secret_manager.rs @@ -12,8 +12,8 @@ use aws_sdk_secretsmanager::{ }; use litellm_secrets_aws::{AwsSecretsManagerV2, Error, RotationResponse}; use litellm_secrets_types::{ - AwsOperationContext, BaseSecretManager, KeyManagementSettings, SecretOperationContext, - SecretValue, SecretWriteContext, + AwsOperationContext, BaseSecretManager, KeyManagementSettings, Secret, SecretDeleter, + SecretValue, SecretWriteContext, SecretWriter, }; use rstest::{fixture, rstest}; use serde_json::json; @@ -22,459 +22,13 @@ use wiremock::{ matchers::{body_partial_json, header}, }; -fn manager(server: &MockServer, settings: KeyManagementSettings) -> AwsSecretsManagerV2 { - let client = Client::from_conf( - aws_sdk_secretsmanager::Config::builder() - .behavior_version(BehaviorVersion::latest()) - .region(Region::new("us-east-1")) - .credentials_provider(Credentials::new("test", "test", None, None, "test")) - .endpoint_url(server.uri()) - .retry_config(RetryConfig::disabled()) - .build(), - ); - AwsSecretsManagerV2::new(client, (&settings).into()) -} +#[path = "secret_manager/support.rs"] +mod support; +use support::*; -fn loaded_manager(server: &MockServer) -> AwsSecretsManagerV2 { - let endpoint_url = server.uri(); - let environment: Arc = - Arc::new(move |name: &str| match name { - "AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(endpoint_url.clone()), - "AWS_ACCESS_KEY_ID" => Some("test".into()), - "AWS_SECRET_ACCESS_KEY" => Some("test".into()), - _ => None, - }); - AwsSecretsManagerV2::load_aws_secret_manager( - Some(true), - KeyManagementSettings { - aws_region_name: Some("us-east-1".into()), - ..Default::default() - }, - environment, - ) - .unwrap() - .unwrap() -} - -#[fixture] -fn default_settings() -> KeyManagementSettings { - KeyManagementSettings::default() -} - -#[rstest] -#[case::string_value("KEY", Some("value"))] -#[case::missing_value("missing", None)] -#[case::non_string_value("BOOL", None)] -#[tokio::test] -async fn primary_lookup_preserves_read_semantics( - default_settings: KeyManagementSettings, - #[case] name: &str, - #[case] expected: Option<&str>, -) { - let server = MockServer::start().await; - Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) - .and(body_partial_json(json!({"SecretId":"primary"}))) - .respond_with( - ResponseTemplate::new(200).set_body_json( - json!({"SecretString":json!({"KEY":"value", "BOOL":true}).to_string()}), - ), - ) - .expect(1) - .mount(&server) - .await; - let manager = manager(&server, default_settings); - assert_eq!( - manager - .read_secret_for_resolver(name, Some("primary"), &|_: &str| None) - .await - .unwrap() - .and_then(|v| v.as_str().map(str::to_owned)) - .as_deref(), - expected - ); -} - -#[rstest] -#[case::access_key("AWS_ACCESS_KEY_ID")] -#[case::secret_access_key("AWS_SECRET_ACCESS_KEY")] -#[case::region_name("AWS_REGION_NAME")] -#[case::region("AWS_REGION")] -#[case::bedrock_endpoint("AWS_BEDROCK_RUNTIME_ENDPOINT")] -#[tokio::test] -async fn bootstrap_keys_bypass_primary_lookup( - default_settings: KeyManagementSettings, - #[case] name: &str, -) { - let server = MockServer::start().await; - let manager = manager(&server, default_settings); - assert_eq!( - manager - .read_secret_for_resolver(name, Some("primary"), &|_: &str| Some("bootstrap".into())) - .await - .unwrap() - .unwrap() - .as_str() - .unwrap(), - "bootstrap" - ); -} - -#[rstest] -#[tokio::test] -async fn failed_read_returns_none_but_invalid_primary_json_is_an_error( - default_settings: KeyManagementSettings, -) { - let server = MockServer::start().await; - Mock::given(body_partial_json(json!({"SecretId":"missing"}))) - .respond_with( - ResponseTemplate::new(400).set_body_json(json!({"__type":"ResourceNotFoundException"})), - ) - .mount(&server) - .await; - Mock::given(body_partial_json(json!({"SecretId":"invalid"}))) - .respond_with(ResponseTemplate::new(200).set_body_json(json!({"SecretString":"not-json"}))) - .mount(&server) - .await; - Mock::given(body_partial_json(json!({"SecretId":"no-string"}))) - .respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name":"no-string"}))) - .mount(&server) - .await; - let manager = manager(&server, default_settings); - assert!( - manager - .async_read_secret("missing") - .await - .unwrap() - .is_none() - ); - assert!(matches!( - manager - .read_secret_for_resolver("KEY", Some("invalid"), &|_: &str| None) - .await, - Err(Error::PrimarySecret) - )); - assert!(matches!( - manager.async_read_secret("no-string").await, - Err(Error::MissingString) - )); -} - -#[rstest] -#[tokio::test] -async fn same_name_rotation_uses_put_and_returns_its_response( - default_settings: KeyManagementSettings, -) { - let server = MockServer::start().await; - Mock::given(header("x-amz-target", "secretsmanager.PutSecretValue")) - .and(body_partial_json( - json!({"SecretId":"key", "SecretString":"replacement"}), - )) - .respond_with( - ResponseTemplate::new(200).set_body_json(json!({"Name":"key", "VersionId":"version"})), - ) - .expect(1) - .mount(&server) - .await; - let response = manager(&server, default_settings) - .async_rotate_secret("key", "key", &SecretValue::new("replacement")) - .await - .unwrap(); - match response { - RotationResponse::Updated(output) => assert_eq!(output.version_id(), Some("version")), - _ => panic!("rotation created a second secret"), - } - assert_eq!(server.received_requests().await.unwrap().len(), 1); -} - -#[rstest] -#[tokio::test] -async fn renamed_rotation_reads_creates_verifies_then_deletes( - default_settings: KeyManagementSettings, -) { - let server = MockServer::start().await; - let step = AtomicUsize::new(0); - Mock::given(wiremock::matchers::method("POST")) - .respond_with(move |request: &wiremock::Request| { - let body: serde_json::Value = request.body_json().unwrap(); - let action = request - .headers - .get("x-amz-target") - .unwrap() - .to_str() - .unwrap(); - match step.fetch_add(1, Ordering::SeqCst) { - 0 => { - assert_eq!(action, "secretsmanager.GetSecretValue"); - assert_eq!(body["SecretId"], "old"); - ResponseTemplate::new(200).set_body_json(json!({"SecretString":"old-value"})) - } - 1 => { - assert_eq!(action, "secretsmanager.CreateSecret"); - assert_eq!(body["Name"], "new"); - assert_eq!(body["Description"], "Rotated from old"); - assert_eq!(body["SecretString"], "replacement"); - ResponseTemplate::new(200).set_body_json(json!({"Name":"new"})) - } - 2 => { - assert_eq!(action, "secretsmanager.GetSecretValue"); - assert_eq!(body["SecretId"], "new"); - ResponseTemplate::new(200).set_body_json(json!({"SecretString":"replacement"})) - } - 3 => { - assert_eq!(action, "secretsmanager.DeleteSecret"); - assert_eq!(body["SecretId"], "old"); - assert_eq!(body["RecoveryWindowInDays"], 7); - ResponseTemplate::new(200).set_body_json(json!({"Name":"old"})) - } - _ => panic!("unexpected request"), - } - }) - .expect(4) - .mount(&server) - .await; - assert!(matches!( - manager(&server, default_settings) - .async_rotate_secret("old", "new", &SecretValue::new("replacement")) - .await - .unwrap(), - RotationResponse::Created(_) - )); -} - -#[rstest] -#[tokio::test] -async fn creation_passes_tags_and_kms_and_survives_replication_failure() { - let server = MockServer::start().await; - Mock::given(header("x-amz-target", "secretsmanager.CreateSecret")) - .and(body_partial_json(json!({"Name":"key", "SecretString":"value", "KmsKeyId":"kms-key", "Tags":[{"Key":"stage", "Value":"test"}]}))) - .respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name":"key"}))).expect(1).mount(&server).await; - Mock::given(header( - "x-amz-target", - "secretsmanager.ReplicateSecretToRegions", - )) - .and(body_partial_json( - json!({"SecretId":"key", "AddReplicaRegions":[{"Region":"replica-region"}]}), - )) - .respond_with( - ResponseTemplate::new(400).set_body_json(json!({"__type":"InvalidRequestException"})), - ) - .expect(1) - .mount(&server) - .await; - let settings = KeyManagementSettings { - kms_key_id: Some("kms-key".into()), - tags: Some(std::collections::BTreeMap::from([( - "stage".into(), - "test".into(), - )])), - replica_regions: Some(vec!["replica-region".into()]), - ..Default::default() - }; - let manager = manager(&server, settings); - assert_eq!( - manager - .async_write_secret("key", &SecretValue::new("value"), None) - .await - .unwrap() - .name(), - Some("key") - ); - assert!( - manager - .async_replicate_secret("key", &[]) - .await - .unwrap() - .is_none() - ); -} - -#[rstest] -#[tokio::test] -async fn trait_write_uses_typed_write_context(default_settings: KeyManagementSettings) { - let server = MockServer::start().await; - Mock::given(header("x-amz-target", "secretsmanager.CreateSecret")) - .and(body_partial_json(json!({ - "Name": "key", - "SecretString": "value", - "Description": "created by caller", - "Tags": [{"Key": "stage", "Value": "test"}], - }))) - .respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name": "key"}))) - .expect(1) - .mount(&server) - .await; - let context = SecretWriteContext { - description: Some("created by caller".into()), - tags: std::collections::BTreeMap::from([("stage".into(), "test".into())]), - ..Default::default() - }; - let response = BaseSecretManager::async_write_secret( - &manager(&server, default_settings), - "key", - &SecretValue::new("value"), - &context, - ) - .await - .unwrap(); - assert_eq!(response.name(), Some("key")); -} - -#[rstest] -#[tokio::test] -async fn trait_delete_accepts_an_unspecified_recovery_window( - default_settings: KeyManagementSettings, -) { - let server = MockServer::start().await; - Mock::given(header("x-amz-target", "secretsmanager.DeleteSecret")) - .and(body_partial_json(json!({"SecretId": "key"}))) - .respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name": "key"}))) - .expect(1) - .mount(&server) - .await; - let response = BaseSecretManager::async_delete_secret( - &manager(&server, default_settings), - "key", - None, - &SecretOperationContext::default(), - ) - .await - .unwrap(); - assert_eq!(response.name(), Some("key")); -} - -#[rstest] -#[tokio::test] -async fn trait_read_uses_the_aws_region_from_its_operation_context() { - let server = MockServer::start().await; - Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) - .respond_with(|request: &wiremock::Request| { - let authorization = request - .headers - .get("authorization") - .unwrap() - .to_str() - .unwrap(); - assert!(authorization.contains("/us-west-2/secretsmanager/aws4_request")); - ResponseTemplate::new(200).set_body_json(json!({"SecretString": "value"})) - }) - .expect(1) - .mount(&server) - .await; - let context = SecretOperationContext::Aws(AwsOperationContext { - region_name: Some("us-west-2".into()), - ..Default::default() - }); - let value = BaseSecretManager::async_read_secret(&loaded_manager(&server), "key", &context) - .await - .unwrap(); - assert_eq!(value.unwrap().expose(), "value"); -} - -#[rstest] -#[tokio::test] -async fn trait_read_applies_the_aws_operation_timeout() { - let server = MockServer::start().await; - Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) - .respond_with( - ResponseTemplate::new(200) - .set_delay(Duration::from_secs(1)) - .set_body_json(json!({"SecretString": "late"})), - ) - .expect(1) - .mount(&server) - .await; - let context = SecretOperationContext::Aws(AwsOperationContext { - timeout: Some(Duration::from_millis(30)), - ..Default::default() - }); - assert!(matches!( - BaseSecretManager::async_read_secret(&loaded_manager(&server), "key", &context).await, - Err(Error::Timeout) - )); -} - -#[rstest] -#[tokio::test] -async fn credential_failures_are_not_swallowed_as_missing_secrets() { - use aws_credential_types::provider::{ProvideCredentials, error::CredentialsError, future}; - #[derive(Debug)] - struct FailedCredentials; - impl ProvideCredentials for FailedCredentials { - fn provide_credentials<'a>(&'a self) -> future::ProvideCredentials<'a> - where - Self: 'a, - { - future::ProvideCredentials::ready(Err(CredentialsError::provider_error( - "private-auth-detail", - ))) - } - } - let server = MockServer::start().await; - let config = aws_sdk_secretsmanager::Config::builder() - .behavior_version(BehaviorVersion::latest()) - .region(Region::new("us-east-1")) - .credentials_provider(FailedCredentials) - .endpoint_url(server.uri()) - .retry_config(RetryConfig::disabled()) - .build(); - let manager = AwsSecretsManagerV2::new(Client::from_conf(config), Default::default()); - let error = manager.async_read_secret("key").await.unwrap_err(); - assert!(!format!("{error:?}").contains("private-auth-detail")); - assert!(matches!(error, Error::Read(_))); - assert!(server.received_requests().await.unwrap().is_empty()); -} - -#[rstest] -#[tokio::test] -async fn read_timeout_is_an_error_and_cannot_be_mistaken_for_missing() { - let server = MockServer::start().await; - Mock::given(wiremock::matchers::method("POST")) - .respond_with( - ResponseTemplate::new(200) - .set_delay(Duration::from_secs(1)) - .set_body_json(json!({"SecretString":"late"})), - ) - .mount(&server) - .await; - let config = aws_sdk_secretsmanager::Config::builder() - .behavior_version(BehaviorVersion::latest()) - .region(Region::new("us-east-1")) - .credentials_provider(Credentials::new("test", "test", None, None, "test")) - .endpoint_url(server.uri()) - .retry_config(RetryConfig::disabled()) - .timeout_config( - aws_sdk_secretsmanager::config::timeout::TimeoutConfig::builder() - .operation_timeout(Duration::from_millis(30)) - .build(), - ) - .build(); - let manager = AwsSecretsManagerV2::new(Client::from_conf(config), Default::default()); - assert!(matches!( - manager.async_read_secret("key").await, - Err(Error::Timeout) - )); -} - -#[rstest] -#[case::denied(400, "AccessDeniedException")] -#[case::throttled(400, "ThrottlingException")] -#[case::unavailable(503, "ServiceUnavailableException")] -#[tokio::test] -async fn service_failures_remain_errors( - default_settings: KeyManagementSettings, - #[case] status: u16, - #[case] code: &str, -) { - let server = MockServer::start().await; - Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) - .respond_with(ResponseTemplate::new(status).set_body_json(json!({"__type":code}))) - .expect(1) - .mount(&server) - .await; - assert!(matches!( - manager(&server, default_settings) - .async_read_secret("key") - .await, - Err(Error::Read(_)) - )); -} +#[path = "secret_manager/configuration.rs"] +mod configuration; +#[path = "secret_manager/reads.rs"] +mod reads; +#[path = "secret_manager/writes.rs"] +mod writes; diff --git a/litellm-rust/crates/secrets-aws/tests/secret_manager/configuration.rs b/litellm-rust/crates/secrets-aws/tests/secret_manager/configuration.rs new file mode 100644 index 00000000000..40b851fd678 --- /dev/null +++ b/litellm-rust/crates/secrets-aws/tests/secret_manager/configuration.rs @@ -0,0 +1,256 @@ +use super::*; + +#[rstest] +#[tokio::test] +async fn credential_failures_are_not_swallowed_as_missing_secrets() { + use aws_credential_types::provider::{ProvideCredentials, error::CredentialsError, future}; + #[derive(Debug)] + struct FailedCredentials; + impl ProvideCredentials for FailedCredentials { + fn provide_credentials<'a>(&'a self) -> future::ProvideCredentials<'a> + where + Self: 'a, + { + future::ProvideCredentials::ready(Err(CredentialsError::provider_error( + "private-auth-detail", + ))) + } + } + let server = MockServer::start().await; + let config = client_builder(&server) + .credentials_provider(FailedCredentials) + .build(); + let manager = AwsSecretsManagerV2::new(Client::from_conf(config), Default::default()); + let error = manager.async_read_secret("key").await.unwrap_err(); + assert!(!format!("{error:?}").contains("private-auth-detail")); + assert!(matches!(error, Error::Read(_))); + assert!(server.received_requests().await.unwrap().is_empty()); +} + +#[rstest] +#[case::environment(false)] +#[case::operation_override(true)] +#[tokio::test] +async fn endpoint_overrides_replace_the_service_and_override_the_region( + #[case] override_context: bool, +) { + let configured = MockServer::start().await; + let explicit = MockServer::start().await; + let target = if override_context { + &explicit + } else { + &configured + }; + Mock::given(wiremock::matchers::path_regex("^/secretsmanager/?$")) + .and(header("x-amz-target", "secretsmanager.GetSecretValue")) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"SecretString":"value"}))) + .expect(1) + .mount(target) + .await; + let endpoint = format!("{}/bedrock-runtime", configured.uri()); + let manager = AwsSecretsManagerV2::load_aws_secret_manager( + Some(true), + KeyManagementSettings { + aws_region_name: Some("cn-north-1".into()), + ..Default::default() + }, + Arc::new(move |name: &str| match name { + "AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(endpoint.clone()), + "AWS_ACCESS_KEY_ID" | "AWS_SECRET_ACCESS_KEY" => Some("test".into()), + _ => None, + }), + ) + .unwrap() + .unwrap(); + let context = AwsOperationContext { + bedrock_runtime_endpoint: override_context + .then(|| format!("{}/bedrock-runtime", explicit.uri())), + ..Default::default() + }; + assert_eq!( + BaseSecretManager::async_read_secret(&manager, "key", &context) + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); + assert!( + if override_context { + configured + } else { + explicit + } + .received_requests() + .await + .unwrap() + .is_empty() + ); +} + +#[rstest] +#[case::unset(None)] +#[case::disabled(Some(false))] +fn disabled_secret_manager_loader_does_not_require_environment(#[case] enabled: Option) { + assert!( + AwsSecretsManagerV2::load_aws_secret_manager( + enabled, + Default::default(), + Arc::new(|_: &str| panic!("disabled loader consulted the environment")) + ) + .unwrap() + .is_none() + ); +} + +#[rstest] +#[case::role(None, None)] +#[case::cross_account(Some("external-id"), None)] +#[case::web_identity(None, Some("identity-token"))] +#[tokio::test] +async fn configured_sts_credentials_sign_the_secret_request( + #[case] external_id: Option<&str>, + #[case] identity: Option<&str>, +) { + use aws_sdk_secretsmanager::primitives::{DateTime, DateTimeFormat}; + let server = MockServer::start().await; + let expiry = DateTime::from(std::time::SystemTime::now() + Duration::from_secs(3600)) + .fmt(DateTimeFormat::DateTime) + .unwrap(); + let action = if identity.is_some() { + "AssumeRoleWithWebIdentity" + } else { + "AssumeRole" + }; + let expected_external = external_id.map(str::to_owned); + let expected_identity = identity.map(str::to_owned); + Mock::given(wiremock::matchers::body_string_contains(format!("Action={action}"))) + .respond_with(move |request: &wiremock::Request| { + let body = std::str::from_utf8(&request.body).unwrap(); + assert!(body.contains("RoleArn=test-role"), "{body}"); + assert!(body.contains("RoleSessionName=parity-session"), "{body}"); + if let Some(value) = &expected_external { assert!(body.contains(&format!("ExternalId={value}"))); } + if let Some(value) = &expected_identity { assert!(body.contains(&format!("WebIdentityToken={value}"))); } + ResponseTemplate::new(200).set_body_string(format!( + "<{action}Response><{action}Result>assumed-key\ + assumed-secretsession-token\ + {expiry}")) + }).expect(1).mount(&server).await; + Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) + .and(header("x-amz-security-token", "session-token")) + .respond_with(|request: &wiremock::Request| { + assert!( + request.headers["authorization"] + .to_str() + .unwrap() + .contains("Credential=assumed-key/") + ); + ResponseTemplate::new(200).set_body_json(json!({"SecretString":"value"})) + }) + .expect(1) + .mount(&server) + .await; + let endpoint = server.uri(); + let manager = AwsSecretsManagerV2::load_aws_secret_manager( + Some(true), + KeyManagementSettings { + aws_region_name: Some("us-east-1".into()), + aws_role_name: Some("test-role".into()), + aws_session_name: Some("parity-session".into()), + aws_external_id: external_id.map(SecretValue::new), + aws_web_identity_token: identity.map(SecretValue::new), + aws_sts_endpoint: Some(endpoint.clone()), + ..Default::default() + }, + Arc::new(move |name: &str| match name { + "AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(endpoint.clone()), + "AWS_ACCESS_KEY_ID" | "AWS_SECRET_ACCESS_KEY" => Some("source-key".into()), + _ => None, + }), + ) + .unwrap() + .unwrap(); + assert_eq!( + manager + .async_read_secret("key") + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); +} + +#[tokio::test] +async fn configured_profile_credentials_override_static_environment_credentials() { + const CHILD_ENDPOINT: &str = "LITELLM_SECRETS_PROFILE_TEST_ENDPOINT"; + if let Ok(endpoint) = std::env::var(CHILD_ENDPOINT) { + let manager = AwsSecretsManagerV2::load_aws_secret_manager( + Some(true), + KeyManagementSettings { + aws_region_name: Some("us-east-1".into()), + aws_profile_name: Some("parity".into()), + ..Default::default() + }, + Arc::new(move |name: &str| match name { + "AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(endpoint.clone()), + "AWS_ACCESS_KEY_ID" | "AWS_SECRET_ACCESS_KEY" => Some("wrong-static-key".into()), + _ => None, + }), + ) + .unwrap() + .unwrap(); + assert_eq!( + manager + .async_read_secret("key") + .await + .unwrap() + .unwrap() + .expose(), + "profile-value" + ); + return; + } + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) + .and(header("x-amz-security-token", "profile-session")) + .respond_with(|request: &wiremock::Request| { + assert!( + request.headers["authorization"] + .to_str() + .unwrap() + .contains("Credential=profile-key/") + ); + ResponseTemplate::new(200).set_body_json(json!({"SecretString":"profile-value"})) + }) + .expect(1) + .mount(&server) + .await; + let directory = tempfile::tempdir().unwrap(); + let credentials = directory.path().join("credentials"); + let config = directory.path().join("config"); + std::fs::write(&credentials, "[parity]\naws_access_key_id=profile-key\naws_secret_access_key=profile-secret\naws_session_token=profile-session\n").unwrap(); + std::fs::write(&config, "").unwrap(); + let endpoint = server.uri(); + let result = tokio::task::spawn_blocking(move || { + std::process::Command::new(std::env::current_exe().unwrap()) + .args([ + "--exact", + "configuration::configured_profile_credentials_override_static_environment_credentials", + "--nocapture", + ]) + .env(CHILD_ENDPOINT, endpoint) + .env("AWS_SHARED_CREDENTIALS_FILE", credentials) + .env("AWS_CONFIG_FILE", config) + .output() + .unwrap() + }) + .await + .unwrap(); + assert!( + result.status.success(), + "{}\n{}", + String::from_utf8_lossy(&result.stdout), + String::from_utf8_lossy(&result.stderr) + ); +} diff --git a/litellm-rust/crates/secrets-aws/tests/secret_manager/reads.rs b/litellm-rust/crates/secrets-aws/tests/secret_manager/reads.rs new file mode 100644 index 00000000000..3c39924f09f --- /dev/null +++ b/litellm-rust/crates/secrets-aws/tests/secret_manager/reads.rs @@ -0,0 +1,198 @@ +use super::*; + +#[rstest] +#[case::string_value("KEY", Some(Secret::String(SecretValue::new("value"))))] +#[case::missing_value("missing", None)] +#[case::non_string_value("BOOL", Some(Secret::Bool(true)))] +#[tokio::test] +async fn primary_lookup_preserves_read_semantics( + default_settings: KeyManagementSettings, + #[case] name: &str, + #[case] expected: Option, +) { + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) + .and(body_partial_json(json!({"SecretId":"primary"}))) + .respond_with( + ResponseTemplate::new(200).set_body_json( + json!({"SecretString":json!({"KEY":"value", "BOOL":true}).to_string()}), + ), + ) + .expect(1) + .mount(&server) + .await; + let manager = manager(&server, default_settings); + assert_eq!( + manager + .read_secret_for_resolver(name, Some("primary"), &|_: &str| None) + .await + .unwrap(), + expected + ); +} + +#[rstest] +#[case::access_key("AWS_ACCESS_KEY_ID")] +#[case::secret_access_key("AWS_SECRET_ACCESS_KEY")] +#[case::region_name("AWS_REGION_NAME")] +#[case::region("AWS_REGION")] +#[case::bedrock_endpoint("AWS_BEDROCK_RUNTIME_ENDPOINT")] +#[tokio::test] +async fn bootstrap_keys_bypass_primary_lookup( + default_settings: KeyManagementSettings, + #[case] name: &str, +) { + let server = MockServer::start().await; + let manager = manager(&server, default_settings); + assert_eq!( + manager + .read_secret_for_resolver(name, Some("primary"), &|_: &str| Some("bootstrap".into())) + .await + .unwrap() + .unwrap() + .as_str() + .unwrap(), + "bootstrap" + ); +} + +#[rstest] +#[tokio::test] +async fn failed_read_returns_none_but_invalid_primary_json_is_an_error( + default_settings: KeyManagementSettings, +) { + let server = MockServer::start().await; + Mock::given(body_partial_json(json!({"SecretId":"missing"}))) + .respond_with( + ResponseTemplate::new(400).set_body_json(json!({"__type":"ResourceNotFoundException"})), + ) + .mount(&server) + .await; + Mock::given(body_partial_json(json!({"SecretId":"invalid"}))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"SecretString":"not-json"}))) + .mount(&server) + .await; + Mock::given(body_partial_json(json!({"SecretId":"no-string"}))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name":"no-string"}))) + .mount(&server) + .await; + let manager = manager(&server, default_settings); + assert!( + manager + .async_read_secret("missing") + .await + .unwrap() + .is_none() + ); + assert!(matches!( + manager + .read_secret_for_resolver("KEY", Some("invalid"), &|_: &str| None) + .await, + Err(Error::PrimarySecret) + )); + assert!(matches!( + manager.async_read_secret("no-string").await, + Err(Error::MissingString) + )); +} + +#[rstest] +#[tokio::test] +async fn trait_read_uses_the_aws_region_from_its_operation_context() { + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) + .respond_with(|request: &wiremock::Request| { + let authorization = request + .headers + .get("authorization") + .unwrap() + .to_str() + .unwrap(); + assert!(authorization.contains("/us-west-2/secretsmanager/aws4_request")); + ResponseTemplate::new(200).set_body_json(json!({"SecretString": "value"})) + }) + .expect(1) + .mount(&server) + .await; + let context = AwsOperationContext { + region_name: Some("us-west-2".into()), + ..Default::default() + }; + let value = BaseSecretManager::async_read_secret(&loaded_manager(&server), "key", &context) + .await + .unwrap(); + assert_eq!(value.unwrap().expose(), "value"); +} + +#[rstest] +#[tokio::test] +async fn trait_read_applies_the_aws_operation_timeout() { + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) + .respond_with( + ResponseTemplate::new(200) + .set_delay(Duration::from_secs(1)) + .set_body_json(json!({"SecretString": "late"})), + ) + .expect(1) + .mount(&server) + .await; + let context = AwsOperationContext { + timeout: Some(Duration::from_millis(30)), + ..Default::default() + }; + assert!(matches!( + BaseSecretManager::async_read_secret(&loaded_manager(&server), "key", &context).await, + Err(Error::Timeout) + )); +} + +#[rstest] +#[tokio::test] +async fn read_timeout_is_an_error_and_cannot_be_mistaken_for_missing() { + let server = MockServer::start().await; + Mock::given(wiremock::matchers::method("POST")) + .respond_with( + ResponseTemplate::new(200) + .set_delay(Duration::from_secs(1)) + .set_body_json(json!({"SecretString":"late"})), + ) + .mount(&server) + .await; + let config = client_builder(&server) + .timeout_config( + aws_sdk_secretsmanager::config::timeout::TimeoutConfig::builder() + .operation_timeout(Duration::from_millis(30)) + .build(), + ) + .build(); + let manager = AwsSecretsManagerV2::new(Client::from_conf(config), Default::default()); + assert!(matches!( + manager.async_read_secret("key").await, + Err(Error::Timeout) + )); +} + +#[rstest] +#[case::denied(400, "AccessDeniedException")] +#[case::throttled(400, "ThrottlingException")] +#[case::unavailable(503, "ServiceUnavailableException")] +#[tokio::test] +async fn service_failures_remain_errors( + default_settings: KeyManagementSettings, + #[case] status: u16, + #[case] code: &str, +) { + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.GetSecretValue")) + .respond_with(ResponseTemplate::new(status).set_body_json(json!({"__type":code}))) + .expect(1) + .mount(&server) + .await; + assert!(matches!( + manager(&server, default_settings) + .async_read_secret("key") + .await, + Err(Error::Read(_)) + )); +} diff --git a/litellm-rust/crates/secrets-aws/tests/secret_manager/support.rs b/litellm-rust/crates/secrets-aws/tests/secret_manager/support.rs new file mode 100644 index 00000000000..6e0df2fdab7 --- /dev/null +++ b/litellm-rust/crates/secrets-aws/tests/secret_manager/support.rs @@ -0,0 +1,80 @@ +use super::*; + +pub(super) fn manager(server: &MockServer, settings: KeyManagementSettings) -> AwsSecretsManagerV2 { + let client = Client::from_conf(client_builder(server).build()); + AwsSecretsManagerV2::new(client, (&settings).into()) +} + +pub(super) fn loaded_manager(server: &MockServer) -> AwsSecretsManagerV2 { + let endpoint_url = server.uri(); + let environment: Arc = + Arc::new(move |name: &str| match name { + "AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(endpoint_url.clone()), + "AWS_ACCESS_KEY_ID" => Some("test".into()), + "AWS_SECRET_ACCESS_KEY" => Some("test".into()), + _ => None, + }); + AwsSecretsManagerV2::load_aws_secret_manager( + Some(true), + KeyManagementSettings { + aws_region_name: Some("us-east-1".into()), + ..Default::default() + }, + environment, + ) + .unwrap() + .unwrap() +} + +#[fixture] +pub(super) fn default_settings() -> KeyManagementSettings { + KeyManagementSettings::default() +} + +pub(super) async fn scripted_actions(server: &MockServer, actions: Vec) { + let count = actions.len() as u64; + let step = AtomicUsize::new(0); + Mock::given(wiremock::matchers::method("POST")) + .respond_with(move |request: &wiremock::Request| { + let Action { + operation: action, + request: expected, + status, + response, + } = &actions[step.fetch_add(1, Ordering::SeqCst)]; + assert_eq!( + request.headers["x-amz-target"], + format!("secretsmanager.{action}") + ); + let body: serde_json::Value = request.body_json().unwrap(); + let actual = serde_json::Value::Object( + body.as_object() + .unwrap() + .iter() + .filter(|(key, _)| key.as_str() != "ClientRequestToken") + .map(|(key, value)| (key.clone(), value.clone())) + .collect(), + ); + assert_eq!(&actual, expected); + ResponseTemplate::new(*status).set_body_json(response) + }) + .expect(count) + .mount(server) + .await; +} + +pub(super) struct Action { + pub(super) operation: &'static str, + pub(super) request: serde_json::Value, + pub(super) status: u16, + pub(super) response: serde_json::Value, +} + +pub(super) fn client_builder(server: &MockServer) -> aws_sdk_secretsmanager::config::Builder { + aws_sdk_secretsmanager::Config::builder() + .behavior_version(BehaviorVersion::latest()) + .region(Region::new("us-east-1")) + .credentials_provider(Credentials::new("test", "test", None, None, "test")) + .endpoint_url(server.uri()) + .retry_config(RetryConfig::disabled()) +} diff --git a/litellm-rust/crates/secrets-aws/tests/secret_manager/writes.rs b/litellm-rust/crates/secrets-aws/tests/secret_manager/writes.rs new file mode 100644 index 00000000000..9968837d549 --- /dev/null +++ b/litellm-rust/crates/secrets-aws/tests/secret_manager/writes.rs @@ -0,0 +1,615 @@ +use super::*; + +#[rstest] +#[tokio::test] +async fn same_name_rotation_uses_put_and_returns_its_response( + default_settings: KeyManagementSettings, +) { + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.PutSecretValue")) + .and(body_partial_json( + json!({"SecretId":"key", "SecretString":"replacement"}), + )) + .respond_with( + ResponseTemplate::new(200).set_body_json(json!({"Name":"key", "VersionId":"version"})), + ) + .expect(1) + .mount(&server) + .await; + let response = manager(&server, default_settings) + .async_rotate_secret("key", "key", &SecretValue::new("replacement")) + .await + .unwrap(); + match response { + RotationResponse::Updated(output) => assert_eq!(output.version_id(), Some("version")), + _ => panic!("rotation created a second secret"), + } + assert_eq!(server.received_requests().await.unwrap().len(), 1); +} + +#[rstest] +#[tokio::test] +async fn renamed_rotation_reads_creates_verifies_then_deletes( + default_settings: KeyManagementSettings, +) { + let server = MockServer::start().await; + scripted_actions( + &server, + vec![ + Action { + operation: "GetSecretValue", + request: json!({"SecretId":"old"}), + status: 200, + response: json!({"SecretString":"old-value"}), + }, + Action { + operation: "CreateSecret", + request: json!({"Name":"new", "Description":"Rotated from old", "SecretString":"replacement"}), + status: 200, + response: json!({"Name":"new"}), + }, + Action { + operation: "GetSecretValue", + request: json!({"SecretId":"new"}), + status: 200, + response: json!({"SecretString":"replacement"}), + }, + Action { + operation: "DeleteSecret", + request: json!({"SecretId":"old", "RecoveryWindowInDays":7}), + status: 200, + response: json!({"Name":"old"}), + }, + ], + ).await; + assert!(matches!( + manager(&server, default_settings) + .async_rotate_secret("old", "new", &SecretValue::new("replacement")) + .await + .unwrap(), + RotationResponse::Created(_) + )); +} + +#[rstest] +#[tokio::test] +async fn creation_passes_tags_and_kms_and_survives_replication_failure() { + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.CreateSecret")) + .and(body_partial_json(json!({"Name":"key", "SecretString":"value", "KmsKeyId":"kms-key", "Tags":[{"Key":"stage", "Value":"test"}]}))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name":"key"}))).expect(1).mount(&server).await; + Mock::given(header( + "x-amz-target", + "secretsmanager.ReplicateSecretToRegions", + )) + .and(body_partial_json( + json!({"SecretId":"key", "AddReplicaRegions":[{"Region":"replica-region"}]}), + )) + .respond_with( + ResponseTemplate::new(400).set_body_json(json!({"__type":"InvalidRequestException"})), + ) + .expect(1) + .mount(&server) + .await; + let settings = KeyManagementSettings { + kms_key_id: Some("kms-key".into()), + tags: Some(std::collections::BTreeMap::from([( + "stage".into(), + "test".into(), + )])), + replica_regions: Some(vec!["replica-region".into()]), + ..Default::default() + }; + let manager = manager(&server, settings); + assert_eq!( + manager + .async_write_secret("key", &SecretValue::new("value"), None) + .await + .unwrap() + .name(), + Some("key") + ); + assert!( + manager + .async_replicate_secret("key", &[]) + .await + .unwrap() + .is_none() + ); +} + +#[rstest] +#[tokio::test] +async fn trait_write_uses_typed_write_context(default_settings: KeyManagementSettings) { + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.CreateSecret")) + .and(body_partial_json(json!({ + "Name": "key", + "SecretString": "value", + "Description": "created by caller", + "Tags": [{"Key": "stage", "Value": "test"}], + }))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name": "key"}))) + .expect(1) + .mount(&server) + .await; + let context = SecretWriteContext { + description: Some("created by caller".into()), + tags: std::collections::BTreeMap::from([("stage".into(), "test".into())]), + ..Default::default() + }; + let response = SecretWriter::async_write_secret( + &manager(&server, default_settings), + "key", + &SecretValue::new("value"), + &context, + ) + .await + .unwrap(); + assert_eq!(response.name(), Some("key")); +} + +#[rstest] +#[tokio::test] +async fn trait_delete_uses_the_provider_recovery_policy(default_settings: KeyManagementSettings) { + let server = MockServer::start().await; + Mock::given(header("x-amz-target", "secretsmanager.DeleteSecret")) + .and(body_partial_json(json!({"SecretId": "key"}))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({"Name": "key"}))) + .expect(1) + .mount(&server) + .await; + let response = SecretDeleter::async_delete_secret( + &manager(&server, default_settings), + "key", + &AwsOperationContext::default(), + ) + .await + .unwrap(); + assert_eq!(response.name(), Some("key")); +} + +#[rstest] +#[case::write(false)] +#[case::rotate_back(true)] +#[tokio::test] +async fn recovery_window_alias_is_restored_updated_and_tagged(#[case] rotate: bool) { + let server = MockServer::start().await; + let description = if rotate { + "Rotated from old" + } else { + "description" + }; + let write = json!({"Name":"key", "SecretString":"new", "Description":description, + "KmsKeyId":"kms", "Tags":[{"Key":"stage", "Value":"test"}]}); + let actions = if rotate { + vec![Action { + operation: "GetSecretValue", + request: json!({"SecretId":"old"}), + status: 200, + response: json!({"SecretString":"old"}), + }] + } else { + vec![] + }; + let recovery = vec![ + Action { + operation: "CreateSecret", + request: write, + status: 400, + response: json!({"__type":"ResourceExistsException"}), + }, + Action { + operation: "DescribeSecret", + request: json!({"SecretId":"key"}), + status: 200, + response: json!({"DeletedDate":1}), + }, + Action { + operation: "RestoreSecret", + request: json!({"SecretId":"key"}), + status: 200, + response: json!({"Name":"key"}), + }, + Action { + operation: "UpdateSecret", + request: json!({"SecretId":"key", "SecretString":"new", "Description":description, + "KmsKeyId":"kms"}), + status: 200, + response: json!({"ARN":"restored-arn", "Name":"key", "VersionId":"new-version"}), + }, + Action { + operation: "TagResource", + request: json!({"SecretId":"key", "Tags":[{"Key":"stage", "Value":"test"}]}), + status: 200, + response: json!({}), + }, + ]; + let verification = if rotate { + vec![ + Action { + operation: "GetSecretValue", + request: json!({"SecretId":"key"}), + status: 200, + response: json!({"SecretString":"new"}), + }, + Action { + operation: "DeleteSecret", + request: json!({"SecretId":"old", "RecoveryWindowInDays":7}), + status: 200, + response: json!({}), + }, + ] + } else { + vec![] + }; + scripted_actions( + &server, + actions + .into_iter() + .chain(recovery) + .chain(verification) + .collect(), + ) + .await; + let manager = manager( + &server, + KeyManagementSettings { + kms_key_id: Some("kms".into()), + tags: Some(std::collections::BTreeMap::from([( + "stage".into(), + "test".into(), + )])), + ..Default::default() + }, + ); + let output = if rotate { + match manager + .async_rotate_secret("old", "key", &SecretValue::new("new")) + .await + .unwrap() + { + RotationResponse::Created(output) => output, + _ => panic!("expected restored alias"), + } + } else { + manager + .async_write_secret("key", &SecretValue::new("new"), Some(description)) + .await + .unwrap() + }; + assert_eq!( + (output.arn(), output.name(), output.version_id()), + (Some("restored-arn"), Some("key"), Some("new-version")) + ); +} + +#[rstest] +#[case::live(200, json!({"Name":"key"}))] +#[case::missing(400, json!({"__type":"ResourceNotFoundException"}))] +#[case::denied(400, json!({"__type":"AccessDeniedException"}))] +#[tokio::test] +async fn create_failure_does_not_overwrite_an_alias_without_a_deletion_date( + #[case] status: u16, + #[case] described: serde_json::Value, +) { + let server = MockServer::start().await; + scripted_actions( + &server, + vec![ + Action { + operation: "CreateSecret", + request: json!({"Name":"key", "SecretString":"new"}), + status: 400, + response: json!({"__type":"ResourceExistsException"}), + }, + Action { + operation: "DescribeSecret", + request: json!({"SecretId":"key"}), + status, + response: described, + }, + ], + ) + .await; + assert!(matches!( + manager(&server, Default::default()) + .async_write_secret("key", &SecretValue::new("new"), None) + .await, + Err(Error::Create(_)) + )); +} + +#[derive(Clone, Copy, Debug)] +enum RecoveryFailure { + Restore, + Update, + Tag, + DeleteAfterUpdate, + DeleteAfterTag, +} + +#[rstest] +#[case::unconfigured(None)] +#[case::empty(Some(vec![]))] +#[case::configured(Some(vec!["region-a".into(), "region-b".into()]))] +#[tokio::test] +async fn creation_replicates_only_to_configured_regions(#[case] regions: Option>) { + let server = MockServer::start().await; + let create = vec![Action { + operation: "CreateSecret", + request: json!({"Name":"key", "SecretString":"value", "KmsKeyId":"kms-key"}), + status: 200, + response: json!({"Name":"key", "VersionId":"created"}), + }]; + let replicate = regions + .as_ref() + .filter(|regions| !regions.is_empty()) + .map(|regions| Action { + operation: "ReplicateSecretToRegions", + request: json!({"SecretId":"key", "AddReplicaRegions":regions.iter() + .map(|region| json!({"Region":region})).collect::>()}), + status: 200, + response: json!({"ARN":"replica-arn"}), + }); + scripted_actions(&server, create.into_iter().chain(replicate).collect()).await; + let environment: Arc = { + let endpoint = server.uri(); + Arc::new(move |name: &str| match name { + "AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(endpoint.clone()), + "AWS_ACCESS_KEY_ID" | "AWS_SECRET_ACCESS_KEY" => Some("test".into()), + _ => None, + }) + }; + let manager = AwsSecretsManagerV2::load_aws_secret_manager( + Some(true), + KeyManagementSettings { + aws_region_name: Some("us-east-1".into()), + replica_regions: regions, + kms_key_id: Some("kms-key".into()), + ..Default::default() + }, + environment, + ) + .unwrap() + .unwrap(); + assert_eq!( + manager + .async_write_secret("key", &SecretValue::new("value"), None) + .await + .unwrap() + .version_id(), + Some("created") + ); +} + +#[rstest] +#[case::success(200)] +#[case::denied(403)] +#[tokio::test] +async fn direct_replication_returns_response_or_service_error(#[case] status: u16) { + let server = MockServer::start().await; + scripted_actions( + &server, + vec![Action { + operation: "ReplicateSecretToRegions", + request: json!({"SecretId":"key", "AddReplicaRegions":[{"Region":"region-a"}, {"Region":"region-b"}]}), + status, + response: if status == 200 { + json!({"ARN":"replicated-arn"}) + } else { + json!({"__type":"AccessDeniedException"}) + }, + }], + ).await; + let result = manager(&server, Default::default()) + .async_replicate_secret("key", &["region-a".into(), "region-b".into()]) + .await; + if status == 200 { + assert_eq!(result.unwrap().unwrap().arn(), Some("replicated-arn")); + } else { + assert!(matches!(result, Err(Error::Replicate(_)))); + } +} + +#[rstest] +#[case::create(false)] +#[case::replicate(true)] +#[tokio::test] +async fn write_and_replication_timeouts_remain_errors(#[case] replicate: bool) { + let server = MockServer::start().await; + Mock::given(wiremock::matchers::method("POST")) + .respond_with( + ResponseTemplate::new(200) + .set_delay(Duration::from_secs(1)) + .set_body_json(json!({})), + ) + .expect(if replicate { 1 } else { 2 }) + .mount(&server) + .await; + let client = Client::from_conf( + client_builder(&server) + .timeout_config( + aws_sdk_secretsmanager::config::timeout::TimeoutConfig::builder() + .operation_timeout(Duration::from_millis(50)) + .build(), + ) + .build(), + ); + let manager = AwsSecretsManagerV2::new(client, Default::default()); + if replicate { + assert!(matches!( + manager + .async_replicate_secret("key", &["region".into()]) + .await, + Err(Error::Replicate(_)) + )); + } else { + assert!(matches!( + manager + .async_write_secret("key", &SecretValue::new("value"), None) + .await, + Err(Error::Create(_)) + )); + } +} + +#[rstest] +#[case::text("value")] +#[case::json(r#"{"api_key":"test","metadata":{"team":"test"},"temperature":0.7}"#)] +#[case::empty("")] +#[case::unicode(" Ï€\n ")] +#[tokio::test] +async fn write_read_delete_preserves_the_complete_secret_string(#[case] value: &str) { + let server = MockServer::start().await; + scripted_actions( + &server, + vec![ + Action { + operation: "CreateSecret", + request: json!({"Name":"key", "SecretString":value, "Description":"description"}), + status: 200, + response: json!({"Name":"key"}), + }, + Action { + operation: "GetSecretValue", + request: json!({"SecretId":"key"}), + status: 200, + response: json!({"SecretString":value}), + }, + Action { + operation: "DeleteSecret", + request: json!({"SecretId":"key", "RecoveryWindowInDays":7}), + status: 200, + response: json!({"Name":"key"}), + }, + ], + ) + .await; + let manager = manager(&server, Default::default()); + assert_eq!( + manager + .async_write_secret("key", &SecretValue::new(value), Some("description")) + .await + .unwrap() + .name(), + Some("key") + ); + assert_eq!( + manager + .async_read_secret("key") + .await + .unwrap() + .unwrap() + .expose(), + value + ); + assert_eq!( + manager + .async_delete_secret("key", Some(7)) + .await + .unwrap() + .name(), + Some("key") + ); +} + +#[rstest] +#[case::restore(RecoveryFailure::Restore)] +#[case::update(RecoveryFailure::Update)] +#[case::tag(RecoveryFailure::Tag)] +#[case::delete_after_update(RecoveryFailure::DeleteAfterUpdate)] +#[case::delete_after_tag(RecoveryFailure::DeleteAfterTag)] +#[tokio::test] +async fn failed_update_reschedules_deletion_of_a_restored_alias(#[case] failure: RecoveryFailure) { + let server = MockServer::start().await; + let response = |failed| { + if failed { + (400, json!({"__type":"InvalidRequestException"})) + } else { + (200, json!({})) + } + }; + let (restore_status, restore_body) = response(matches!(failure, RecoveryFailure::Restore)); + let (update_status, update_body) = response(matches!( + failure, + RecoveryFailure::Update | RecoveryFailure::DeleteAfterUpdate + )); + let (delete_status, delete_body) = response(matches!( + failure, + RecoveryFailure::DeleteAfterUpdate | RecoveryFailure::DeleteAfterTag + )); + let prefix = [ + Action { + operation: "CreateSecret", + request: json!({"Name":"key", "SecretString":"new", "Tags":[{"Key":"stage", "Value":"test"}]}), + status: 400, + response: json!({"__type":"ResourceExistsException"}), + }, + Action { + operation: "DescribeSecret", + request: json!({"SecretId":"key"}), + status: 200, + response: json!({"DeletedDate":1}), + }, + Action { + operation: "RestoreSecret", + request: json!({"SecretId":"key"}), + status: restore_status, + response: restore_body, + }, + ]; + let update = (!matches!(failure, RecoveryFailure::Restore)).then_some(Action { + operation: "UpdateSecret", + request: json!({"SecretId":"key", "SecretString":"new"}), + status: update_status, + response: update_body, + }); + let tag = matches!( + failure, + RecoveryFailure::Tag | RecoveryFailure::DeleteAfterTag + ) + .then_some(Action { + operation: "TagResource", + request: json!({"SecretId":"key", "Tags":[{"Key":"stage", "Value":"test"}]}), + status: 400, + response: json!({"__type":"InvalidRequestException"}), + }); + let delete = (!matches!(failure, RecoveryFailure::Restore)).then_some(Action { + operation: "DeleteSecret", + request: json!({"SecretId":"key", "RecoveryWindowInDays":7}), + status: delete_status, + response: delete_body, + }); + scripted_actions( + &server, + prefix + .into_iter() + .chain(update) + .chain(tag) + .chain(delete) + .collect(), + ) + .await; + let error = manager( + &server, + KeyManagementSettings { + tags: Some(std::collections::BTreeMap::from([( + "stage".into(), + "test".into(), + )])), + ..Default::default() + }, + ) + .async_write_secret("key", &SecretValue::new("new"), None) + .await + .unwrap_err(); + match failure { + RecoveryFailure::Restore => assert!(matches!(error, Error::Restore(_))), + RecoveryFailure::Update => assert!(matches!(error, Error::Update(_))), + RecoveryFailure::Tag => assert!(matches!(error, Error::Tag(_))), + RecoveryFailure::DeleteAfterUpdate | RecoveryFailure::DeleteAfterTag => { + assert!(matches!(error, Error::Delete(_))) + } + } +} diff --git a/litellm-rust/crates/secrets-azure/AGENTS.md b/litellm-rust/crates/secrets-azure/AGENTS.md new file mode 100644 index 00000000000..8fcd32a6c0c --- /dev/null +++ b/litellm-rust/crates/secrets-azure/AGENTS.md @@ -0,0 +1 @@ +- https://learn.microsoft.com/en-us/rest/api/keyvault/secrets/get-secret/get-secret diff --git a/litellm-rust/crates/secrets-azure/Cargo.toml b/litellm-rust/crates/secrets-azure/Cargo.toml index 96db7f235ef..7e8a79f89ef 100644 --- a/litellm-rust/crates/secrets-azure/Cargo.toml +++ b/litellm-rust/crates/secrets-azure/Cargo.toml @@ -6,6 +6,7 @@ license.workspace = true repository.workspace = true [dependencies] +tokio.workspace = true litellm-auth-azure.workspace = true litellm-auth-types.workspace = true litellm-secrets-types.workspace = true @@ -17,7 +18,6 @@ veil.workspace = true percent-encoding = "2.3" [dev-dependencies] -tokio.workspace = true wiremock = "0.6.5" rstest.workspace = true serde_json.workspace = true diff --git a/litellm-rust/crates/secrets-azure/src/error.rs b/litellm-rust/crates/secrets-azure/src/error.rs index 9b20efe4f7c..2862f08a56c 100644 --- a/litellm-rust/crates/secrets-azure/src/error.rs +++ b/litellm-rust/crates/secrets-azure/src/error.rs @@ -1,5 +1,9 @@ #[derive(thiserror::Error, veil::Redact)] pub enum Error { + #[error(transparent)] + Operation(#[from] litellm_secrets_types::Error), + #[error("secret manager operation timed out")] + Timeout, #[error("{0} environment variable is missing")] MissingEnvironment(&'static str), #[error("AZURE_KEY_VAULT_URI is not a valid https vault URL")] diff --git a/litellm-rust/crates/secrets-azure/src/key_vault.rs b/litellm-rust/crates/secrets-azure/src/key_vault.rs index 81ad0331ca7..13e59f8e4ac 100644 --- a/litellm-rust/crates/secrets-azure/src/key_vault.rs +++ b/litellm-rust/crates/secrets-azure/src/key_vault.rs @@ -3,7 +3,7 @@ use std::sync::Arc; use litellm_auth_azure::{AzureAuthInputs, AzureAuthService, ConfigValue}; use litellm_auth_types::{InputSource, Sourced}; use litellm_core_utils::settings::Lookup; -use litellm_secrets_types::{Secret, SecretValue}; +use litellm_secrets_types::{AzureOperationContext, BaseSecretManager, Secret, SecretValue}; use percent_encoding::{AsciiSet, NON_ALPHANUMERIC}; use serde::Deserialize; @@ -77,6 +77,12 @@ impl AzureKeyVault { } pub async fn get_secret(&self, name: &str) -> Result, Error> { + BaseSecretManager::async_read_secret(self, name, &AzureOperationContext::default()) + .await + .map(|value| value.map(Secret::String)) + } + + async fn read(&self, name: &str) -> Result, Error> { let token = self .auth .get_azure_ad_token(&self.inputs, &|key| self.environment.get(key)) @@ -91,6 +97,7 @@ impl AzureKeyVault { .client .get(url) .bearer_auth(token.value().secret().expose()) + .header(reqwest::header::ACCEPT, "application/json") .send() .await .map_err(Error::Http)?; @@ -102,7 +109,7 @@ impl AzureKeyVault { } let payload: SecretResponse = response.json().await.map_err(Error::Http)?; let value = payload.value.ok_or(Error::MissingValue)?; - Ok(Some(Secret::String(SecretValue::new(value)))) + Ok(Some(SecretValue::new(value))) } } @@ -113,3 +120,59 @@ fn scope_for(vault: &reqwest::Url) -> String { .map_or(host, |(_, remainder)| remainder); format!("https://{resource}/.default") } + +impl BaseSecretManager for AzureKeyVault { + type Error = Error; + type Context = AzureOperationContext; + + async fn async_read_secret( + &self, + name: &str, + context: &Self::Context, + ) -> Result, Error> { + match context.timeout { + Some(timeout) => tokio::time::timeout(timeout, self.read(name)) + .await + .map_err(|_| Error::Timeout)?, + None => self.read(name).await, + } + } +} + +pub trait AzureTokenProvider: Send + Sync { + fn get_token<'a>( + &'a self, + scope: &'a str, + environment: &'a (dyn Lookup + Send + Sync), + ) -> std::pin::Pin> + Send + 'a>>; +} + +#[derive(Default)] +pub struct NativeAzureTokenProvider { + auth: AzureAuthService, +} + +impl AzureTokenProvider for NativeAzureTokenProvider { + fn get_token<'a>( + &'a self, + scope: &'a str, + environment: &'a (dyn Lookup + Send + Sync), + ) -> std::pin::Pin> + Send + 'a>> + { + Box::pin(async move { + let inputs = AzureAuthInputs { + azure_scope: ConfigValue::Value(Sourced::new( + scope.to_owned(), + InputSource::Deployment, + )), + enable_azure_ad_token_refresh: Sourced::new(true, InputSource::Deployment), + ..Default::default() + }; + self.auth + .get_azure_ad_token(&inputs, &|name| environment.get(name)) + .await? + .map(|token| SecretValue::new(token.value().secret().expose())) + .ok_or(Error::MissingCredentials) + }) + } +} diff --git a/litellm-rust/crates/secrets-azure/src/lib.rs b/litellm-rust/crates/secrets-azure/src/lib.rs index c0094fc033b..8c1a182db6b 100644 --- a/litellm-rust/crates/secrets-azure/src/lib.rs +++ b/litellm-rust/crates/secrets-azure/src/lib.rs @@ -4,4 +4,4 @@ mod error; mod key_vault; pub use error::Error; -pub use key_vault::AzureKeyVault; +pub use key_vault::{AzureKeyVault, AzureTokenProvider, NativeAzureTokenProvider}; diff --git a/litellm-rust/crates/secrets-azure/tests/key_vault.rs b/litellm-rust/crates/secrets-azure/tests/key_vault.rs index 3b199f81264..a21149db345 100644 --- a/litellm-rust/crates/secrets-azure/tests/key_vault.rs +++ b/litellm-rust/crates/secrets-azure/tests/key_vault.rs @@ -25,6 +25,7 @@ async fn reads_secret_with_bearer_token_and_api_version() { Mock::given(path("/secrets/OPENAI-API-KEY")) .and(query_param("api-version", "7.4")) .and(header("authorization", "Bearer fake")) + .and(header("accept", "application/json")) .respond_with( ResponseTemplate::new(200) .set_body_json(serde_json::json!({"value": "s3cret", "id": "secret-id"})), @@ -42,6 +43,23 @@ async fn reads_secret_with_bearer_token_and_api_version() { assert_eq!(secret, Secret::String(SecretValue::new("s3cret"))); } +#[rstest] +#[tokio::test] +async fn preserves_secret_contents_and_redacts_debug_output() { + let server = MockServer::start().await; + let value = " \tvalue-Ï€\n"; + Mock::given(path("/secrets/NAME")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({"value": value}))) + .expect(1) + .mount(&server) + .await; + + let secret = manager(&server).get_secret("NAME").await.unwrap().unwrap(); + + assert_eq!(secret.as_str(), Some(value)); + assert!(!format!("{secret:?}").contains(value)); +} + #[rstest] #[tokio::test] async fn percent_encodes_secret_name_path_segment() { @@ -230,3 +248,22 @@ async fn parity_fixture_matches_python_backend_contract(parity_fixture: Fixture) } } } + +#[tokio::test] +async fn trait_read_limits_the_operation_duration() { + use litellm_secrets_types::{AzureOperationContext, BaseSecretManager}; + use std::time::Duration; + let server = MockServer::start().await; + Mock::given(wiremock::matchers::method("GET")) + .respond_with(ResponseTemplate::new(200).set_delay(Duration::from_secs(1))) + .mount(&server) + .await; + let manager = manager(&server); + let context = AzureOperationContext { + timeout: Some(Duration::from_millis(30)), + }; + assert!(matches!( + BaseSecretManager::async_read_secret(&manager, "key", &context).await, + Err(Error::Timeout) + )); +} diff --git a/litellm-rust/crates/secrets-cyberark/AGENTS.md b/litellm-rust/crates/secrets-cyberark/AGENTS.md new file mode 100644 index 00000000000..608850729c4 --- /dev/null +++ b/litellm-rust/crates/secrets-cyberark/AGENTS.md @@ -0,0 +1 @@ +- https://docs.cyberark.com/conjur-open-source/latest/en/content/developer/conjur_api_retrieve_secret.htm diff --git a/litellm-rust/crates/secrets-cyberark/Cargo.toml b/litellm-rust/crates/secrets-cyberark/Cargo.toml index c1613f96403..1c280171f4c 100644 --- a/litellm-rust/crates/secrets-cyberark/Cargo.toml +++ b/litellm-rust/crates/secrets-cyberark/Cargo.toml @@ -19,7 +19,9 @@ percent-encoding = "2.3" tokio = { workspace = true, features = ["sync"] } [dev-dependencies] +rcgen = "0.14.10" rstest.workspace = true +tempfile = "3.27.0" tokio.workspace = true wiremock = "0.6.5" serde.workspace = true diff --git a/litellm-rust/crates/secrets-cyberark/src/error.rs b/litellm-rust/crates/secrets-cyberark/src/error.rs index 5a14f4f3db8..3dfeb95fe26 100644 --- a/litellm-rust/crates/secrets-cyberark/src/error.rs +++ b/litellm-rust/crates/secrets-cyberark/src/error.rs @@ -1,5 +1,7 @@ #[derive(thiserror::Error, veil::Redact)] pub enum Error { + #[error("CyberArk Conjur operation timed out")] + Timeout, #[error("CyberArk Conjur HTTP request failed")] Http( #[from] diff --git a/litellm-rust/crates/secrets-cyberark/src/lib.rs b/litellm-rust/crates/secrets-cyberark/src/lib.rs index 5288f8116b1..74a8b6febf2 100644 --- a/litellm-rust/crates/secrets-cyberark/src/lib.rs +++ b/litellm-rust/crates/secrets-cyberark/src/lib.rs @@ -4,4 +4,4 @@ mod error; mod secret_manager; pub use error::Error; -pub use secret_manager::{CyberArkSecretManager, DeleteOutcome}; +pub use secret_manager::{AuthenticationRetry, CyberArkSecretManager, DeleteOutcome, WriteFailure}; diff --git a/litellm-rust/crates/secrets-cyberark/src/secret_manager.rs b/litellm-rust/crates/secrets-cyberark/src/secret_manager.rs index dd36cf183d8..252a99c917f 100644 --- a/litellm-rust/crates/secrets-cyberark/src/secret_manager.rs +++ b/litellm-rust/crates/secrets-cyberark/src/secret_manager.rs @@ -1,9 +1,14 @@ +mod client; +mod read; +mod write; + use std::{fs, sync::Arc, time::Duration}; use base64::{Engine, engine::general_purpose::STANDARD}; use litellm_core_utils::settings::Lookup; use litellm_secrets_types::{ - BaseSecretManager, SecretOperationContext, SecretValue, SecretWriteContext, + BaseSecretManager, CyberarkOperationContext, RotationError, SecretCache, SecretDeleter, + SecretRotator, SecretValue, SecretWriteContext, SecretWriter, async_rotate_secret, validate_secret_name, }; use moka::future::Cache; @@ -23,6 +28,7 @@ const DEFAULT_API_BASE: &str = "http://127.0.0.1:8080"; const DEFAULT_ACCOUNT: &str = "default"; const DEFAULT_USERNAME: &str = "admin"; const DEFAULT_REFRESH_INTERVAL: Duration = Duration::from_secs(300); +const MAX_TOKEN_LIFETIME: Duration = Duration::from_secs(7 * 60); const SECRET_NAME_SAFE: &AsciiSet = &NON_ALPHANUMERIC .remove(b'-') .remove(b'_') @@ -37,7 +43,7 @@ pub struct CyberArkSecretManager { username: String, api_key: SecretValue, token: Cache<(), SecretValue>, - secrets: Cache, + secrets: SecretCache, authentication_lock: Arc>, } @@ -46,346 +52,53 @@ pub enum DeleteOutcome { NotSupported, } -impl CyberArkSecretManager { - pub fn with_client( - client: reqwest::Client, - endpoint: reqwest::Url, - account: String, - username: String, - api_key: SecretValue, - refresh_interval: Option, - ) -> Self { - let endpoint = normalize_endpoint(endpoint); - let ttl = refresh_interval - .filter(|interval| !interval.is_zero()) - .unwrap_or(DEFAULT_REFRESH_INTERVAL); - let token = Cache::builder().time_to_live(ttl).build(); - let secrets = Cache::builder().time_to_live(ttl).build(); +#[derive(Clone, Copy)] +pub enum AuthenticationRetry { + Never, + Unauthorized, +} + +#[derive(veil::Redact)] +pub struct WriteFailure { + pub source: Error, + #[redact] + pub request_url: Option, + pub authentication: bool, +} + +impl WriteFailure { + fn local(source: Error) -> Self { Self { - client, - endpoint, - account, - username, - api_key, - token, - secrets, - authentication_lock: Arc::new(tokio::sync::Mutex::new(())), + source, + request_url: None, + authentication: false, } } - pub fn new( - environment: Arc, - enterprise_enabled: bool, - ) -> Result { - let api_key = environment.get(CYBERARK_API_KEY).unwrap_or_default(); - let cert = environment.get(CYBERARK_CLIENT_CERT).unwrap_or_default(); - let key = environment.get(CYBERARK_CLIENT_KEY).unwrap_or_default(); - if api_key.is_empty() && (cert.is_empty() || key.is_empty()) { - return Err(Error::MissingCredentials); + fn request(source: Error, url: reqwest::Url) -> Self { + Self { + source, + request_url: Some(url), + authentication: false, } - if !enterprise_enabled { - return Err(Error::EnterpriseRequired); - } - let verify = environment - .get(CYBERARK_SSL_VERIFY) - .map(|value| !value.trim().eq_ignore_ascii_case("false")) - .unwrap_or(true); - let mut builder = reqwest::Client::builder(); - if !verify { - litellm_tracing::warn!( - "CyberArk SSL verification is disabled. This is insecure and should only be used for testing with self-signed certificates." - ); - builder = builder.danger_accept_invalid_certs(true); - } - if !cert.is_empty() && !key.is_empty() { - let certificate = fs::read(cert).map_err(|_| Error::ClientCertificate)?; - let private_key = fs::read(key).map_err(|_| Error::ClientCertificate)?; - let identity = reqwest::Identity::from_pem(&[certificate, private_key].concat()) - .map_err(|_| Error::ClientCertificate)?; - builder = builder.identity(identity); - } - let client = builder.build()?; - let endpoint = reqwest::Url::parse( - &environment - .get(CYBERARK_API_BASE) - .unwrap_or_else(|| DEFAULT_API_BASE.to_owned()), - ) - .map_err(|_| Error::Endpoint)?; - let account = environment - .get(CYBERARK_ACCOUNT) - .unwrap_or_else(|| DEFAULT_ACCOUNT.to_owned()); - let username = environment - .get(CYBERARK_USERNAME) - .unwrap_or_else(|| DEFAULT_USERNAME.to_owned()); - let refresh_interval = environment - .get(CYBERARK_REFRESH_INTERVAL) - .map(|value| { - value - .parse::() - .map(Duration::from_secs) - .map_err(|_| Error::RefreshInterval) - }) - .transpose()?; - Ok(Self::with_client( - client, - endpoint, - account, - username, - SecretValue::new(api_key), - refresh_interval, - )) } +} +impl CyberArkSecretManager { fn secret_url(&self, name: &str) -> Result { let encoded = utf8_percent_encode(name, SECRET_NAME_SAFE); self.endpoint .join(&format!("secrets/{}/variable/{}", self.account, encoded)) .map_err(|_| Error::Endpoint) } - - async fn authenticate(&self, context: &SecretOperationContext) -> Result { - if let Some(token) = self.token.get(&()).await { - return Ok(token); - } - let _guard = self.authentication_lock.lock().await; - if let Some(token) = self.token.get(&()).await { - return Ok(token); - } - let url = self - .endpoint - .join(&format!( - "authn/{}/{}/authenticate", - self.account, self.username - )) - .map_err(|_| Error::Endpoint)?; - let response = with_timeout( - self.client.post(url).body(self.api_key.expose().to_owned()), - context, - ) - .send() - .await?; - if !response.status().is_success() { - return Err(Error::AuthStatus(response.status().as_u16())); - } - let token = SecretValue::new(STANDARD.encode(response.text().await?)); - self.token.insert((), token.clone()).await; - Ok(token) - } - - async fn authorization_header( - &self, - context: &SecretOperationContext, - ) -> Result { - Ok(format!( - "Token token=\"{}\"", - self.authenticate(context).await?.expose() - )) - } - - pub async fn async_read_secret(&self, name: &str) -> Result, Error> { - self.async_read_secret_with_context(name, &SecretOperationContext::default()) - .await - } - - pub async fn async_read_secret_with_context( - &self, - name: &str, - context: &SecretOperationContext, - ) -> Result, Error> { - if let Some(value) = self.secrets.get(name).await { - return Ok(Some(value)); - } - let response = with_timeout( - self.client - .get(self.secret_url(name)?) - .header("Authorization", self.authorization_header(context).await?), - context, - ) - .send() - .await?; - if response.status() == reqwest::StatusCode::NOT_FOUND { - return Ok(None); - } - if !response.status().is_success() { - return Err(Error::Status(response.status().as_u16())); - } - let value = SecretValue::new(response.text().await?); - self.secrets.insert(name.to_owned(), value.clone()).await; - Ok(Some(value)) - } - - pub async fn async_write_secret( - &self, - name: &str, - value: &SecretValue, - description: Option<&str>, - ) -> Result<(), Error> { - self.async_write_secret_with_context( - name, - value, - description, - &SecretOperationContext::default(), - ) - .await - } - - pub async fn async_write_secret_with_context( - &self, - name: &str, - value: &SecretValue, - _description: Option<&str>, - context: &SecretOperationContext, - ) -> Result<(), Error> { - validate_secret_name(name)?; - self.ensure_variable_exists(name, context).await; - let response = with_timeout( - self.client - .post(self.secret_url(name)?) - .header("Authorization", self.authorization_header(context).await?) - .body(value.expose().to_owned()), - context, - ) - .send() - .await?; - if !response.status().is_success() { - return Err(Error::Status(response.status().as_u16())); - } - self.secrets.insert(name.to_owned(), value.clone()).await; - Ok(()) - } - - async fn ensure_variable_exists(&self, name: &str, context: &SecretOperationContext) { - let policy_url = self - .endpoint - .join(&format!("policies/{}/policy/root", self.account)); - let Ok(policy_url) = policy_url else { - litellm_tracing::warn!("Could not build CyberArk policy endpoint"); - return; - }; - let Ok(authorization) = self.authorization_header(context).await else { - litellm_tracing::warn!( - "Could not authenticate while ensuring CyberArk variable exists" - ); - return; - }; - let body = format!( - "- !variable {}\n", - serde_json::to_string(name).expect("serializing a string cannot fail") - ); - let response = with_timeout( - self.client - .post(policy_url) - .header("Authorization", authorization) - .header("Content-Type", "application/x-yaml") - .body(body), - context, - ) - .send() - .await; - match response { - Ok(response) if response.status().is_success() => {} - Ok(response) - if matches!( - response.status(), - reqwest::StatusCode::CONFLICT | reqwest::StatusCode::UNPROCESSABLE_ENTITY - ) => - { - litellm_tracing::debug!( - "CyberArk variable policy already exists or conflicts: {}", - response.status() - ); - } - Ok(response) => { - litellm_tracing::warn!( - "Could not ensure CyberArk variable exists: {}", - response.status() - ); - } - Err(error) => { - litellm_tracing::warn!("Error ensuring CyberArk variable exists: {error}"); - } - } - } - - pub async fn async_delete_secret( - &self, - name: &str, - recovery_window_in_days: Option, - ) -> Result { - self.async_delete_secret_with_context( - name, - recovery_window_in_days, - &SecretOperationContext::default(), - ) - .await - } - - pub async fn async_delete_secret_with_context( - &self, - name: &str, - _recovery_window_in_days: Option, - _context: &SecretOperationContext, - ) -> Result { - litellm_tracing::warn!( - "CyberArk Conjur does not support direct secret deletion. Secrets must be removed through policy updates." - ); - self.secrets.invalidate(name).await; - Ok(DeleteOutcome::NotSupported) - } -} - -impl BaseSecretManager for CyberArkSecretManager { - type Error = Error; - type WriteResponse = (); - type DeleteResponse = DeleteOutcome; - - async fn async_read_secret( - &self, - name: &str, - context: &SecretOperationContext, - ) -> Result, Error> { - self.async_read_secret_with_context(name, context).await - } - - async fn async_write_secret( - &self, - name: &str, - value: &SecretValue, - context: &SecretWriteContext, - ) -> Result<(), Error> { - self.async_write_secret_with_context( - name, - value, - context.description.as_deref(), - &context.operation, - ) - .await - } - - async fn async_delete_secret( - &self, - name: &str, - recovery_window_in_days: Option, - context: &SecretOperationContext, - ) -> Result { - self.async_delete_secret_with_context(name, recovery_window_in_days, context) - .await - } } fn with_timeout( request: reqwest::RequestBuilder, - context: &SecretOperationContext, + context: &CyberarkOperationContext, ) -> reqwest::RequestBuilder { - match context.timeout() { + match context.timeout { Some(timeout) => request.timeout(timeout), None => request, } } - -fn normalize_endpoint(mut endpoint: reqwest::Url) -> reqwest::Url { - if !endpoint.path().ends_with('/') { - endpoint.set_path(&format!("{}/", endpoint.path())); - } - endpoint -} diff --git a/litellm-rust/crates/secrets-cyberark/src/secret_manager/client.rs b/litellm-rust/crates/secrets-cyberark/src/secret_manager/client.rs new file mode 100644 index 00000000000..1d99fe474d5 --- /dev/null +++ b/litellm-rust/crates/secrets-cyberark/src/secret_manager/client.rs @@ -0,0 +1,147 @@ +use super::*; + +impl CyberArkSecretManager { + pub fn with_client( + client: reqwest::Client, + endpoint: reqwest::Url, + account: String, + username: String, + api_key: SecretValue, + refresh_interval: Option, + ) -> Self { + let endpoint = normalize_endpoint(endpoint); + let ttl = refresh_interval + .filter(|interval| !interval.is_zero()) + .unwrap_or(DEFAULT_REFRESH_INTERVAL); + let token = Cache::builder() + .time_to_live(ttl.min(MAX_TOKEN_LIFETIME)) + .build(); + let secrets = SecretCache::new(200, ttl); + Self { + client, + endpoint, + account, + username, + api_key, + token, + secrets, + authentication_lock: Arc::new(tokio::sync::Mutex::new(())), + } + } + + pub fn new( + environment: Arc, + enterprise_enabled: bool, + ) -> Result { + let api_key = environment.get(CYBERARK_API_KEY).unwrap_or_default(); + let cert = environment.get(CYBERARK_CLIENT_CERT).unwrap_or_default(); + let key = environment.get(CYBERARK_CLIENT_KEY).unwrap_or_default(); + if api_key.is_empty() && (cert.is_empty() || key.is_empty()) { + return Err(Error::MissingCredentials); + } + if !enterprise_enabled { + return Err(Error::EnterpriseRequired); + } + let verify = environment + .get(CYBERARK_SSL_VERIFY) + .map(|value| !value.trim().eq_ignore_ascii_case("false")) + .unwrap_or(true); + let mut builder = reqwest::Client::builder(); + if !verify { + litellm_tracing::warn!( + "CyberArk SSL verification is disabled. This is insecure and should only be used for testing with self-signed certificates." + ); + builder = builder.danger_accept_invalid_certs(true); + } + if !cert.is_empty() && !key.is_empty() { + let certificate = fs::read(cert).map_err(|_| Error::ClientCertificate)?; + let private_key = fs::read(key).map_err(|_| Error::ClientCertificate)?; + let identity = reqwest::Identity::from_pem(&[certificate, private_key].concat()) + .map_err(|_| Error::ClientCertificate)?; + builder = builder.identity(identity); + } + let client = builder.build()?; + let endpoint = reqwest::Url::parse( + &environment + .get(CYBERARK_API_BASE) + .unwrap_or_else(|| DEFAULT_API_BASE.to_owned()), + ) + .map_err(|_| Error::Endpoint)?; + let account = environment + .get(CYBERARK_ACCOUNT) + .unwrap_or_else(|| DEFAULT_ACCOUNT.to_owned()); + let username = environment + .get(CYBERARK_USERNAME) + .unwrap_or_else(|| DEFAULT_USERNAME.to_owned()); + let refresh_interval = environment + .get(CYBERARK_REFRESH_INTERVAL) + .map(|value| { + value + .parse::() + .map(Duration::from_secs) + .map_err(|_| Error::RefreshInterval) + }) + .transpose()?; + Ok(Self::with_client( + client, + endpoint, + account, + username, + SecretValue::new(api_key), + refresh_interval, + )) + } + + pub(super) fn authentication_url(&self) -> Result { + self.endpoint + .join(&format!( + "authn/{}/{}/authenticate", + self.account, + utf8_percent_encode(&self.username, SECRET_NAME_SAFE) + )) + .map_err(|_| Error::Endpoint) + } + + pub(super) async fn authenticate( + &self, + context: &CyberarkOperationContext, + ) -> Result { + if let Some(token) = self.token.get(&()).await { + return Ok(token); + } + let _guard = self.authentication_lock.lock().await; + if let Some(token) = self.token.get(&()).await { + return Ok(token); + } + let url = self.authentication_url()?; + let response = with_timeout( + self.client.post(url).body(self.api_key.expose().to_owned()), + context, + ) + .send() + .await?; + if !response.status().is_success() { + return Err(Error::AuthStatus(response.status().as_u16())); + } + let token = SecretValue::new(STANDARD.encode(response.text().await?)); + self.token.insert((), token.clone()).await; + Ok(token) + } + + pub(super) async fn authorization_header( + &self, + context: &CyberarkOperationContext, + ) -> Result { + Ok(format!( + "Token token=\"{}\"", + self.authenticate(context).await?.expose() + )) + } +} + +fn normalize_endpoint(mut endpoint: reqwest::Url) -> reqwest::Url { + if !endpoint.path().ends_with('/') { + endpoint.set_path(&format!("{}/", endpoint.path())); + } + endpoint +} diff --git a/litellm-rust/crates/secrets-cyberark/src/secret_manager/read.rs b/litellm-rust/crates/secrets-cyberark/src/secret_manager/read.rs new file mode 100644 index 00000000000..daa9ed1114d --- /dev/null +++ b/litellm-rust/crates/secrets-cyberark/src/secret_manager/read.rs @@ -0,0 +1,118 @@ +use super::*; + +impl CyberArkSecretManager { + pub async fn async_read_secret(&self, name: &str) -> Result, Error> { + self.async_read_secret_with_context(name, &CyberarkOperationContext::default()) + .await + } + + pub async fn async_read_secret_with_context( + &self, + name: &str, + context: &CyberarkOperationContext, + ) -> Result, Error> { + self.read_with_retry(name, context, AuthenticationRetry::Unauthorized) + .await + } + + pub async fn read_with_retry( + &self, + name: &str, + context: &CyberarkOperationContext, + retry: AuthenticationRetry, + ) -> Result, Error> { + validate_secret_name(name)?; + let read = self.secrets.read( + name.to_owned(), + self.read_uncached_with_retry(name, context, retry), + ); + match context.timeout { + Some(timeout) => tokio::time::timeout(timeout, read) + .await + .map_err(|_| Error::Timeout)?, + None => read.await, + } + } + + pub(super) async fn read_uncached( + &self, + name: &str, + context: &CyberarkOperationContext, + ) -> Result, Error> { + self.read_uncached_with_retry(name, context, AuthenticationRetry::Unauthorized) + .await + } + + pub async fn read_fresh_with_retry( + &self, + name: &str, + context: &CyberarkOperationContext, + retry: AuthenticationRetry, + ) -> Result, Error> { + validate_secret_name(name)?; + self.secrets + .refresh( + name.to_owned(), + self.read_uncached_with_retry(name, context, retry), + ) + .await + } + + pub async fn invalidate_cached_secret(&self, name: &str) { + self.secrets.invalidate(&name.to_owned()).await; + } + + pub(super) async fn read_uncached_with_retry( + &self, + name: &str, + context: &CyberarkOperationContext, + retry: AuthenticationRetry, + ) -> Result, Error> { + let had_cached_token = self.token.get(&()).await.is_some(); + let response = with_timeout( + self.client + .get(self.secret_url(name)?) + .header("Authorization", self.authorization_header(context).await?), + context, + ) + .send() + .await?; + let response = if matches!(retry, AuthenticationRetry::Unauthorized) + && had_cached_token + && response.status() == reqwest::StatusCode::UNAUTHORIZED + { + self.token.invalidate(&()).await; + with_timeout( + self.client + .get(self.secret_url(name)?) + .header("Authorization", self.authorization_header(context).await?), + context, + ) + .send() + .await? + } else { + response + }; + if response.status() == reqwest::StatusCode::NOT_FOUND { + return Ok(None); + } + if !response.status().is_success() { + return Err(Error::Status(response.status().as_u16())); + } + let value = SecretValue::new(response.text().await?); + Ok(Some(value)) + } +} + +impl BaseSecretManager for CyberArkSecretManager { + type Error = Error; + type Context = CyberarkOperationContext; + + async fn async_read_secret( + &self, + name: &str, + context: &Self::Context, + ) -> Result, Error> { + self.async_read_secret_with_context(name, context).await + } +} diff --git a/litellm-rust/crates/secrets-cyberark/src/secret_manager/write.rs b/litellm-rust/crates/secrets-cyberark/src/secret_manager/write.rs new file mode 100644 index 00000000000..265c5fc6b28 --- /dev/null +++ b/litellm-rust/crates/secrets-cyberark/src/secret_manager/write.rs @@ -0,0 +1,262 @@ +use super::*; + +impl CyberArkSecretManager { + pub async fn async_write_secret( + &self, + name: &str, + value: &SecretValue, + description: Option<&str>, + ) -> Result<(), Error> { + self.async_write_secret_with_context( + name, + value, + description, + &CyberarkOperationContext::default(), + ) + .await + } + + pub async fn async_write_secret_with_context( + &self, + name: &str, + value: &SecretValue, + _description: Option<&str>, + context: &CyberarkOperationContext, + ) -> Result<(), Error> { + self.write_with_retry(name, value, context, AuthenticationRetry::Unauthorized) + .await + .map_err(|failure| failure.source) + } + + pub async fn write_with_retry( + &self, + name: &str, + value: &SecretValue, + context: &CyberarkOperationContext, + retry: AuthenticationRetry, + ) -> Result<(), WriteFailure> { + validate_secret_name(name).map_err(|source| WriteFailure::local(source.into()))?; + self.ensure_variable_exists(name, context).await; + let url = self.secret_url(name).map_err(WriteFailure::local)?; + let response = self.post_value(&url, value, context).await?; + let response = if matches!(retry, AuthenticationRetry::Unauthorized) + && response.status() == reqwest::StatusCode::UNAUTHORIZED + { + self.token.invalidate(&()).await; + self.post_value(&url, value, context).await? + } else { + response + }; + if !response.status().is_success() { + return Err(WriteFailure::request( + Error::Status(response.status().as_u16()), + url, + )); + } + self.secrets.insert(name.to_owned(), value.clone()).await; + Ok(()) + } + + async fn post_value( + &self, + url: &reqwest::Url, + value: &SecretValue, + context: &CyberarkOperationContext, + ) -> Result { + let authorization = + self.authorization_header(context) + .await + .map_err(|source| WriteFailure { + source, + request_url: self.authentication_url().ok(), + authentication: true, + })?; + with_timeout( + self.client + .post(url.clone()) + .header("Authorization", authorization) + .body(value.expose().to_owned()), + context, + ) + .send() + .await + .map_err(|source| WriteFailure::request(source.into(), url.clone())) + } + + pub(super) async fn ensure_variable_exists( + &self, + name: &str, + context: &CyberarkOperationContext, + ) { + let policy_url = self + .endpoint + .join(&format!("policies/{}/policy/root", self.account)); + let Ok(policy_url) = policy_url else { + litellm_tracing::warn!("Could not build CyberArk policy endpoint"); + return; + }; + let Ok(authorization) = self.authorization_header(context).await else { + litellm_tracing::warn!( + "Could not authenticate while ensuring CyberArk variable exists" + ); + return; + }; + let body = format!( + "- !variable {}\n", + serde_json::to_string(name).expect("serializing a string cannot fail") + ); + let response = with_timeout( + self.client + .post(policy_url) + .header("Authorization", authorization) + .header("Content-Type", "application/x-yaml") + .body(body), + context, + ) + .send() + .await; + match response { + Ok(response) if response.status().is_success() => {} + Ok(response) + if matches!( + response.status(), + reqwest::StatusCode::CONFLICT | reqwest::StatusCode::UNPROCESSABLE_ENTITY + ) => + { + litellm_tracing::debug!( + "CyberArk variable policy already exists or conflicts: {}", + response.status() + ); + } + Ok(response) => { + litellm_tracing::warn!( + "Could not ensure CyberArk variable exists: {}", + response.status() + ); + } + Err(error) => { + litellm_tracing::warn!("Error ensuring CyberArk variable exists: {error}"); + } + } + } + + pub async fn async_rotate_secret( + &self, + current_name: &str, + new_name: &str, + value: &SecretValue, + ) -> Result<(), RotationError<(), Error>> { + self.async_rotate_secret_with_context( + current_name, + new_name, + value, + &CyberarkOperationContext::default(), + ) + .await + } + + pub async fn async_rotate_secret_with_context( + &self, + current_name: &str, + new_name: &str, + value: &SecretValue, + context: &CyberarkOperationContext, + ) -> Result<(), RotationError<(), Error>> { + async_rotate_secret(self, current_name, new_name, value, context).await + } + + pub async fn async_delete_secret( + &self, + name: &str, + recovery_window_in_days: Option, + ) -> Result { + self.async_delete_secret_with_context( + name, + recovery_window_in_days, + &CyberarkOperationContext::default(), + ) + .await + } + + pub async fn async_delete_secret_with_context( + &self, + name: &str, + _recovery_window_in_days: Option, + _context: &CyberarkOperationContext, + ) -> Result { + litellm_tracing::warn!( + "CyberArk Conjur does not support direct secret deletion. Secrets must be removed through policy updates." + ); + self.secrets.invalidate(&name.to_owned()).await; + Ok(DeleteOutcome::NotSupported) + } +} + +impl SecretWriter for CyberArkSecretManager { + type WriteResponse = (); + + async fn async_write_secret( + &self, + name: &str, + value: &SecretValue, + context: &SecretWriteContext, + ) -> Result<(), Error> { + self.async_write_secret_with_context( + name, + value, + context.description.as_deref(), + &context.operation, + ) + .await + } +} + +impl SecretDeleter for CyberArkSecretManager { + type DeleteResponse = DeleteOutcome; + + async fn async_delete_secret( + &self, + name: &str, + context: &Self::Context, + ) -> Result { + self.async_delete_secret_with_context(name, None, context) + .await + } +} + +impl SecretRotator for CyberArkSecretManager { + type RotationResponse = (); + + async fn async_read_secret_fresh( + &self, + name: &str, + context: &Self::Context, + ) -> Result, Error> { + validate_secret_name(name)?; + let read = self + .secrets + .refresh(name.to_owned(), self.read_uncached(name, context)); + match context.timeout { + Some(timeout) => tokio::time::timeout(timeout, read) + .await + .map_err(|_| Error::Timeout)?, + None => read.await, + } + } + + async fn async_write_replacement( + &self, + current_name: &str, + new_name: &str, + value: &SecretValue, + context: &Self::Context, + ) -> Result<(), Error> { + SecretWriter::async_write_secret( + self, + new_name, + value, + &SecretWriteContext::rotated_from(current_name, context.clone()), + ) + .await + } +} diff --git a/litellm-rust/crates/secrets-cyberark/tests/secret_manager.rs b/litellm-rust/crates/secrets-cyberark/tests/secret_manager.rs index 59964d90d50..2048e067b6e 100644 --- a/litellm-rust/crates/secrets-cyberark/tests/secret_manager.rs +++ b/litellm-rust/crates/secrets-cyberark/tests/secret_manager.rs @@ -1,10 +1,14 @@ -use std::{sync::Arc, time::Duration}; +use std::{ + sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }, + time::Duration, +}; use base64::{Engine, engine::general_purpose::STANDARD}; use litellm_secrets_cyberark::{CyberArkSecretManager, DeleteOutcome, Error}; -use litellm_secrets_types::{ - BaseSecretManager, CyberarkOperationContext, SecretOperationContext, SecretValue, -}; +use litellm_secrets_types::{BaseSecretManager, CyberarkOperationContext, SecretValue}; use rstest::{fixture, rstest}; use serde::Deserialize; use wiremock::{ @@ -12,546 +16,13 @@ use wiremock::{ matchers::{body_string, header, method, path}, }; -const TOKEN_JSON: &str = r#"{"protected":"p","payload":"q","signature":"s"}"#; +#[path = "secret_manager/support.rs"] +mod support; +use support::*; -#[derive(Deserialize)] -struct ParityFixture { - endpoint: String, - account: String, - username: String, - api_key: String, - authenticate_path: String, - token_json: String, - authorization_header: String, - policy_path: String, - secrets: Vec, -} - -#[derive(Deserialize)] -struct ParitySecret { - name: String, - path: String, - policy_body: String, -} - -#[derive(Debug)] -struct RawPath(String); - -impl Match for RawPath { - fn matches(&self, request: &Request) -> bool { - request.url.path() == self.0 - } -} - -#[fixture] -fn parity_fixture() -> ParityFixture { - serde_json::from_str(include_str!("fixtures/parity.json")).unwrap() -} - -fn manager(server: &MockServer, ttl: Duration) -> CyberArkSecretManager { - CyberArkSecretManager::with_client( - reqwest::Client::new(), - server.uri().parse().unwrap(), - "acct".into(), - "admin".into(), - SecretValue::new("k3y"), - Some(ttl), - ) -} - -async fn mount_auth(server: &MockServer, expected: u64) { - Mock::given(method("POST")) - .and(path("/authn/acct/admin/authenticate")) - .and(body_string("k3y")) - .respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON)) - .expect(expected) - .mount(server) - .await; -} - -#[rstest] -#[tokio::test] -async fn successful_reads_cache_auth_secret_and_redact_values() { - let server = MockServer::start().await; - mount_auth(&server, 1).await; - let token = STANDARD.encode(TOKEN_JSON); - Mock::given(path("/secrets/acct/variable/OPENAI_API_KEY")) - .and(header("authorization", format!("Token token=\"{token}\""))) - .respond_with(ResponseTemplate::new(200).set_body_string("sk-live")) - .expect(1) - .mount(&server) - .await; - let manager = manager(&server, Duration::from_secs(60)); - - for _ in 0..2 { - let value = manager - .async_read_secret("OPENAI_API_KEY") - .await - .unwrap() - .unwrap(); - assert_eq!(value.expose(), "sk-live"); - assert!(!format!("{value:?}").contains("sk-live")); - } -} - -#[rstest] -#[tokio::test] -async fn concurrent_reads_share_authentication_request() { - let server = MockServer::start().await; - Mock::given(path("/authn/acct/admin/authenticate")) - .and(body_string("k3y")) - .respond_with( - ResponseTemplate::new(200) - .set_body_string(TOKEN_JSON) - .set_delay(Duration::from_millis(20)), - ) - .expect(1) - .mount(&server) - .await; - Mock::given(path("/secrets/acct/variable/key")) - .and(header( - "authorization", - format!("Token token=\"{}\"", STANDARD.encode(TOKEN_JSON)), - )) - .respond_with(ResponseTemplate::new(200).set_body_string("value")) - .expect(2) - .mount(&server) - .await; - let manager = manager(&server, Duration::from_secs(60)); - - let (first, second) = tokio::join!( - manager.async_read_secret("key"), - manager.async_read_secret("key") - ); - - assert_eq!(first.unwrap().unwrap().expose(), "value"); - assert_eq!(second.unwrap().unwrap().expose(), "value"); -} - -#[rstest] -#[case::not_found(404)] -#[case::unauthorized(401)] -#[case::forbidden(403)] -#[case::server_error(500)] -#[tokio::test] -async fn failed_reads_are_not_cached(#[case] status: u16) { - let server = MockServer::start().await; - mount_auth(&server, 1).await; - let failing = Mock::given(path("/secrets/acct/variable/key")) - .respond_with(ResponseTemplate::new(status)) - .expect(1) - .mount_as_scoped(&server) - .await; - let manager = manager(&server, Duration::from_secs(60)); - let result = manager.async_read_secret("key").await; - if status == 404 { - assert_eq!(result.unwrap(), None); - } else { - assert!(matches!(result, Err(Error::Status(actual)) if actual == status)); - } - drop(failing); - Mock::given(path("/secrets/acct/variable/key")) - .respond_with(ResponseTemplate::new(200).set_body_string("recovered")) - .expect(1) - .mount(&server) - .await; - for _ in 0..2 { - assert_eq!( - manager - .async_read_secret("key") - .await - .unwrap() - .unwrap() - .expose(), - "recovered" - ); - } -} - -#[rstest] -#[tokio::test] -async fn failed_authentication_is_not_cached_and_does_not_read_secret() { - let server = MockServer::start().await; - let failing = Mock::given(path("/authn/acct/admin/authenticate")) - .respond_with(ResponseTemplate::new(401)) - .expect(1) - .mount_as_scoped(&server) - .await; - let unused_secret = Mock::given(path("/secrets/acct/variable/key")) - .respond_with(ResponseTemplate::new(200).set_body_string("value")) - .expect(0) - .mount_as_scoped(&server) - .await; - let manager = manager(&server, Duration::from_secs(60)); - assert!(matches!( - manager.async_read_secret("key").await, - Err(Error::AuthStatus(401)) - )); - drop(unused_secret); - drop(failing); - mount_auth(&server, 1).await; - Mock::given(path("/secrets/acct/variable/key")) - .respond_with(ResponseTemplate::new(200).set_body_string("value")) - .expect(1) - .mount(&server) - .await; - assert_eq!( - manager - .async_read_secret("key") - .await - .unwrap() - .unwrap() - .expose(), - "value" - ); -} - -#[rstest] -#[tokio::test] -async fn trait_read_applies_cyberark_operation_timeout_to_authentication() { - let server = MockServer::start().await; - Mock::given(path("/authn/acct/admin/authenticate")) - .respond_with( - ResponseTemplate::new(200) - .set_body_string(TOKEN_JSON) - .set_delay(Duration::from_millis(50)), - ) - .mount(&server) - .await; - let manager = manager(&server, Duration::from_secs(60)); - let context = SecretOperationContext::Cyberark(CyberarkOperationContext { - timeout: Some(Duration::from_millis(10)), - }); - - let result = BaseSecretManager::async_read_secret(&manager, "key", &context).await; - - assert!(matches!(result, Err(Error::Http(error)) if error.is_timeout())); -} - -#[rstest] -#[tokio::test] -async fn expired_tokens_and_secrets_are_fetched_again() { - let server = MockServer::start().await; - mount_auth(&server, 2).await; - Mock::given(path("/secrets/acct/variable/key")) - .respond_with(ResponseTemplate::new(200).set_body_string("value")) - .expect(2) - .mount(&server) - .await; - let manager = manager(&server, Duration::from_millis(1)); - for _ in 0..2 { - assert!(manager.async_read_secret("key").await.unwrap().is_some()); - tokio::time::sleep(Duration::from_millis(5)).await; - } -} - -#[rstest] -#[case::plain("OPENAI_API_KEY")] -#[case::path("team/app/key")] -#[case::punctuation("a b+c.d-e_f~g")] -#[case::quote("needs \"quote\"")] -#[tokio::test] -async fn secret_names_use_python_quote_encoding(parity_fixture: ParityFixture, #[case] name: &str) { - let secret = parity_fixture - .secrets - .iter() - .find(|secret| secret.name == name) - .unwrap(); - let server = MockServer::start().await; - mount_auth(&server, 1).await; - Mock::given(RawPath(secret.path.clone())) - .respond_with(ResponseTemplate::new(200).set_body_string("value")) - .expect(1) - .mount(&server) - .await; - assert_eq!( - manager(&server, Duration::from_secs(60)) - .async_read_secret(name) - .await - .unwrap() - .unwrap() - .expose(), - "value" - ); -} - -#[rstest] -#[case::created(201)] -#[case::already_exists(409)] -#[case::unprocessable(422)] -#[case::server_error(500)] -#[tokio::test] -async fn writes_tolerate_policy_status_and_cache_value(#[case] policy_status: u16) { - let server = MockServer::start().await; - mount_auth(&server, 1).await; - Mock::given(path("/policies/acct/policy/root")) - .and(header("content-type", "application/x-yaml")) - .and(body_string("- !variable \"team/app\"\n")) - .respond_with(ResponseTemplate::new(policy_status)) - .expect(1) - .mount(&server) - .await; - Mock::given(path("/secrets/acct/variable/team%2Fapp")) - .and(body_string("v")) - .respond_with(ResponseTemplate::new(200)) - .expect(1) - .mount(&server) - .await; - let manager = manager(&server, Duration::from_secs(60)); - manager - .async_write_secret("team/app", &SecretValue::new("v"), None) - .await - .unwrap(); - assert_eq!( - manager - .async_read_secret("team/app") - .await - .unwrap() - .unwrap() - .expose(), - "v" - ); -} - -#[rstest] -#[tokio::test] -async fn failed_value_write_is_not_cached() { - let server = MockServer::start().await; - mount_auth(&server, 1).await; - Mock::given(path("/policies/acct/policy/root")) - .respond_with(ResponseTemplate::new(409)) - .mount(&server) - .await; - Mock::given(path("/secrets/acct/variable/key")) - .and(body_string("v")) - .respond_with(ResponseTemplate::new(403)) - .expect(1) - .mount(&server) - .await; - Mock::given(path("/secrets/acct/variable/key")) - .respond_with(ResponseTemplate::new(200).set_body_string("recovered")) - .expect(1) - .mount(&server) - .await; - let manager = manager(&server, Duration::from_secs(60)); - assert!(matches!( - manager - .async_write_secret("key", &SecretValue::new("v"), None) - .await, - Err(Error::Status(403)) - )); - assert_eq!( - manager - .async_read_secret("key") - .await - .unwrap() - .unwrap() - .expose(), - "recovered" - ); -} - -#[rstest] -#[case::parent("../etc")] -#[case::embedded_parent("team/../etc")] -#[case::control("key\n")] -#[tokio::test] -async fn unsafe_names_fail_before_http_calls(#[case] name: &str) { - let server = MockServer::start().await; - let manager = manager(&server, Duration::from_secs(60)); - assert!(matches!( - manager - .async_write_secret(name, &SecretValue::new("v"), None) - .await, - Err(Error::Operation( - litellm_secrets_types::Error::UnsafeSecretName - )) - )); -} - -#[rstest] -#[tokio::test] -async fn delete_invalidates_cache_and_reports_not_supported() { - let server = MockServer::start().await; - mount_auth(&server, 1).await; - Mock::given(path("/secrets/acct/variable/key")) - .respond_with(ResponseTemplate::new(200).set_body_string("v")) - .expect(2) - .mount(&server) - .await; - let manager = manager(&server, Duration::from_secs(60)); - assert_eq!( - manager - .async_read_secret("key") - .await - .unwrap() - .unwrap() - .expose(), - "v" - ); - assert_eq!( - manager.async_delete_secret("key", Some(7)).await.unwrap(), - DeleteOutcome::NotSupported - ); - assert_eq!( - manager - .async_read_secret("key") - .await - .unwrap() - .unwrap() - .expose(), - "v" - ); -} - -#[rstest] -fn new_validates_credentials_before_license_and_configuration() { - let empty: Arc = - Arc::new(|_: &str| None); - assert!(matches!( - CyberArkSecretManager::new(empty, true), - Err(Error::MissingCredentials) - )); - assert!(matches!( - CyberArkSecretManager::new( - Arc::new(|name: &str| (name == "CYBERARK_API_KEY").then(|| "k3y".into())), - false - ), - Err(Error::EnterpriseRequired) - )); - assert!(matches!( - CyberArkSecretManager::new( - Arc::new(|name: &str| (name == "CYBERARK_CLIENT_CERT").then(|| "cert".into())), - true - ), - Err(Error::MissingCredentials) - )); - assert!(matches!( - CyberArkSecretManager::new( - Arc::new(|name: &str| match name { - "CYBERARK_API_KEY" => Some("k3y".into()), - "CYBERARK_REFRESH_INTERVAL" => Some("abc".into()), - _ => None, - }), - true - ), - Err(Error::RefreshInterval) - )); - assert!(matches!( - CyberArkSecretManager::new( - Arc::new(|name: &str| match name { - "CYBERARK_API_KEY" => Some("k3y".into()), - "CYBERARK_API_BASE" => Some("not a url".into()), - _ => None, - }), - true - ), - Err(Error::Endpoint) - )); -} - -#[rstest] -#[tokio::test] -async fn new_reads_environment_defaults_end_to_end() { - let server = MockServer::start().await; - Mock::given(path("/authn/default/admin/authenticate")) - .and(body_string("k3y")) - .respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON)) - .mount(&server) - .await; - Mock::given(path("/secrets/default/variable/key")) - .respond_with(ResponseTemplate::new(200).set_body_string("value")) - .mount(&server) - .await; - let endpoint = server.uri(); - let manager = CyberArkSecretManager::new( - Arc::new(move |name: &str| match name { - "CYBERARK_API_BASE" => Some(endpoint.clone()), - "CYBERARK_API_KEY" => Some("k3y".into()), - _ => None, - }), - true, - ) - .unwrap(); - assert_eq!( - manager - .async_read_secret("key") - .await - .unwrap() - .unwrap() - .expose(), - "value" - ); -} - -#[rstest] -fn new_reports_missing_client_certificate_files() { - assert!(matches!( - CyberArkSecretManager::new( - Arc::new(|name: &str| match name { - "CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()), - "CYBERARK_CLIENT_KEY" => Some("/missing/key".into()), - _ => None, - }), - true - ), - Err(Error::ClientCertificate) - )); -} - -#[rstest] -#[tokio::test] -async fn trailing_slash_endpoint_preserves_base_path() { - let server = MockServer::start().await; - Mock::given(path("/prefix/authn/acct/admin/authenticate")) - .and(body_string("k3y")) - .respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON)) - .expect(1) - .mount(&server) - .await; - Mock::given(path("/prefix/secrets/acct/variable/key")) - .respond_with(ResponseTemplate::new(200).set_body_string("value")) - .mount(&server) - .await; - let endpoint = format!("{}/prefix/", server.uri()).parse().unwrap(); - let manager = CyberArkSecretManager::with_client( - reqwest::Client::new(), - endpoint, - "acct".into(), - "admin".into(), - SecretValue::new("k3y"), - Some(Duration::from_secs(60)), - ); - assert_eq!( - manager - .async_read_secret("key") - .await - .unwrap() - .unwrap() - .expose(), - "value" - ); -} - -#[rstest] -fn parity_fixture_matches_authentication_contract(parity_fixture: ParityFixture) { - assert_eq!(parity_fixture.endpoint, "http://conjur.test:8080"); - assert_eq!(parity_fixture.account, "acct"); - assert_eq!(parity_fixture.username, "admin"); - assert_eq!(parity_fixture.api_key, "k3y"); - assert_eq!( - parity_fixture.authenticate_path, - "/authn/acct/admin/authenticate" - ); - assert_eq!(parity_fixture.token_json, TOKEN_JSON); - assert_eq!( - parity_fixture.authorization_header, - format!("Token token=\"{}\"", STANDARD.encode(TOKEN_JSON)) - ); - assert_eq!(parity_fixture.policy_path, "/policies/acct/policy/root"); - assert_eq!(parity_fixture.secrets.len(), 4); - assert_eq!( - parity_fixture.secrets[1].policy_body, - "- !variable \"team/app/key\"\n" - ); -} +#[path = "secret_manager/configuration.rs"] +mod configuration; +#[path = "secret_manager/reads.rs"] +mod reads; +#[path = "secret_manager/writes.rs"] +mod writes; diff --git a/litellm-rust/crates/secrets-cyberark/tests/secret_manager/configuration.rs b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/configuration.rs new file mode 100644 index 00000000000..fbd4317f446 --- /dev/null +++ b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/configuration.rs @@ -0,0 +1,452 @@ +use super::*; + +#[rstest] +#[tokio::test] +async fn successful_reads_cache_auth_secret_and_redact_values() { + let server = MockServer::start().await; + mount_auth(&server, 1).await; + let token = STANDARD.encode(TOKEN_JSON); + Mock::given(path("/secrets/acct/variable/OPENAI_API_KEY")) + .and(header("authorization", format!("Token token=\"{token}\""))) + .respond_with(ResponseTemplate::new(200).set_body_string("sk-live")) + .expect(1) + .mount(&server) + .await; + let manager = manager(&server, Duration::from_secs(60)); + + for _ in 0..2 { + let value = manager + .async_read_secret("OPENAI_API_KEY") + .await + .unwrap() + .unwrap(); + assert_eq!(value.expose(), "sk-live"); + assert!(!format!("{value:?}").contains("sk-live")); + } +} + +#[rstest] +#[tokio::test] +async fn concurrent_reads_share_authentication_and_secret_requests() { + let server = MockServer::start().await; + Mock::given(path("/authn/acct/admin/authenticate")) + .and(body_string("k3y")) + .respond_with( + ResponseTemplate::new(200) + .set_body_string(TOKEN_JSON) + .set_delay(Duration::from_millis(20)), + ) + .expect(1) + .mount(&server) + .await; + Mock::given(path("/secrets/acct/variable/key")) + .and(header( + "authorization", + format!("Token token=\"{}\"", STANDARD.encode(TOKEN_JSON)), + )) + .respond_with(ResponseTemplate::new(200).set_body_string("value")) + .expect(1) + .mount(&server) + .await; + let manager = manager(&server, Duration::from_secs(60)); + + let (first, second) = tokio::join!( + manager.async_read_secret("key"), + manager.async_read_secret("key") + ); + + assert_eq!(first.unwrap().unwrap().expose(), "value"); + assert_eq!(second.unwrap().unwrap().expose(), "value"); +} + +#[rstest] +#[case::host("host/team/app", "/authn/acct/host%2Fteam%2Fapp/authenticate")] +#[case::user("alice@devops", "/authn/acct/alice%40devops/authenticate")] +#[tokio::test] +async fn authentication_encodes_login(#[case] username: &str, #[case] expected_path: &str) { + let server = MockServer::start().await; + Mock::given(RawPath(expected_path.to_owned())) + .and(method("POST")) + .and(body_string("k3y")) + .respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON)) + .expect(1) + .mount(&server) + .await; + Mock::given(path("/secrets/acct/variable/key")) + .respond_with(ResponseTemplate::new(200).set_body_string("value")) + .mount(&server) + .await; + let manager = CyberArkSecretManager::with_client( + reqwest::Client::new(), + server.uri().parse().unwrap(), + "acct".into(), + username.into(), + SecretValue::new("k3y"), + None, + ); + + assert_eq!( + manager + .async_read_secret("key") + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); +} + +#[rstest] +#[tokio::test] +async fn rejected_cached_token_is_reauthenticated_once() { + let server = MockServer::start().await; + mount_auth(&server, 2).await; + Mock::given(path("/secrets/acct/variable/first")) + .respond_with(ResponseTemplate::new(200).set_body_string("first-value")) + .expect(1) + .mount(&server) + .await; + let attempts = Arc::new(AtomicUsize::new(0)); + let attempts_for_response = Arc::clone(&attempts); + Mock::given(path("/secrets/acct/variable/second")) + .respond_with(move |_: &Request| { + if attempts_for_response.fetch_add(1, Ordering::SeqCst) == 0 { + ResponseTemplate::new(401) + } else { + ResponseTemplate::new(200).set_body_string("second-value") + } + }) + .expect(2) + .mount(&server) + .await; + let manager = manager(&server, Duration::from_secs(600)); + + assert_eq!( + manager + .async_read_secret("first") + .await + .unwrap() + .unwrap() + .expose(), + "first-value" + ); + assert_eq!( + manager + .async_read_secret("second") + .await + .unwrap() + .unwrap() + .expose(), + "second-value" + ); + assert_eq!(attempts.load(Ordering::SeqCst), 2); +} + +#[rstest] +#[tokio::test] +async fn failed_authentication_is_not_cached_and_does_not_read_secret() { + let server = MockServer::start().await; + let failing = Mock::given(path("/authn/acct/admin/authenticate")) + .respond_with(ResponseTemplate::new(401)) + .expect(1) + .mount_as_scoped(&server) + .await; + let unused_secret = Mock::given(path("/secrets/acct/variable/key")) + .respond_with(ResponseTemplate::new(200).set_body_string("value")) + .expect(0) + .mount_as_scoped(&server) + .await; + let manager = manager(&server, Duration::from_secs(60)); + assert!(matches!( + manager.async_read_secret("key").await, + Err(Error::AuthStatus(401)) + )); + drop(unused_secret); + drop(failing); + mount_auth(&server, 1).await; + Mock::given(path("/secrets/acct/variable/key")) + .respond_with(ResponseTemplate::new(200).set_body_string("value")) + .expect(1) + .mount(&server) + .await; + assert_eq!( + manager + .async_read_secret("key") + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); +} + +#[rstest] +#[tokio::test] +async fn trait_read_applies_cyberark_operation_timeout_to_authentication() { + let server = MockServer::start().await; + Mock::given(path("/authn/acct/admin/authenticate")) + .respond_with( + ResponseTemplate::new(200) + .set_body_string(TOKEN_JSON) + .set_delay(Duration::from_millis(50)), + ) + .mount(&server) + .await; + let manager = manager(&server, Duration::from_secs(60)); + let context = CyberarkOperationContext { + timeout: Some(Duration::from_millis(10)), + }; + + let result = BaseSecretManager::async_read_secret(&manager, "key", &context).await; + + match result { + Err(Error::Timeout) => {} + Err(Error::Http(error)) => assert!(error.is_timeout()), + other => panic!("expected timeout, got {other:?}"), + } +} + +#[rstest] +fn new_validates_credentials_before_license_and_configuration() { + let empty: Arc = + Arc::new(|_: &str| None); + assert!(matches!( + CyberArkSecretManager::new(empty, true), + Err(Error::MissingCredentials) + )); + assert!(matches!( + CyberArkSecretManager::new( + Arc::new(|name: &str| (name == "CYBERARK_API_KEY").then(|| "k3y".into())), + false + ), + Err(Error::EnterpriseRequired) + )); + assert!(matches!( + CyberArkSecretManager::new( + Arc::new(|name: &str| (name == "CYBERARK_CLIENT_CERT").then(|| "cert".into())), + true + ), + Err(Error::MissingCredentials) + )); + assert!(matches!( + CyberArkSecretManager::new( + Arc::new(|name: &str| match name { + "CYBERARK_API_KEY" => Some("k3y".into()), + "CYBERARK_REFRESH_INTERVAL" => Some("abc".into()), + _ => None, + }), + true + ), + Err(Error::RefreshInterval) + )); + assert!(matches!( + CyberArkSecretManager::new( + Arc::new(|name: &str| match name { + "CYBERARK_API_KEY" => Some("k3y".into()), + "CYBERARK_API_BASE" => Some("not a url".into()), + _ => None, + }), + true + ), + Err(Error::Endpoint) + )); +} + +#[rstest] +fn certificate_only_credentials_are_validated_as_a_client_identity() { + let result = CyberArkSecretManager::new( + Arc::new(|name: &str| match name { + "CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()), + "CYBERARK_CLIENT_KEY" => Some("/missing/key".into()), + _ => None, + }), + true, + ); + + assert!(matches!(result, Err(Error::ClientCertificate))); +} + +#[rstest] +#[case::certificate_only("")] +#[case::certificate_and_api_key("k3y")] +#[tokio::test] +async fn configured_client_identity_preserves_auth_request_and_read_result( + client_identity_directory: tempfile::TempDir, + #[case] api_key: &'static str, +) { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/authn/default/admin/authenticate")) + .and(body_string(api_key)) + .respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON)) + .expect(1) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/secrets/default/variable/key")) + .and(header( + "authorization", + format!("Token token=\"{}\"", STANDARD.encode(TOKEN_JSON)), + )) + .respond_with(ResponseTemplate::new(200).set_body_string(" value\n")) + .expect(1) + .mount(&server) + .await; + let endpoint = server.uri(); + let certificate = client_identity_directory.path().join("client.crt"); + let key = client_identity_directory.path().join("client.key"); + let manager = CyberArkSecretManager::new( + Arc::new(move |name: &str| match name { + "CYBERARK_API_BASE" => Some(endpoint.clone()), + "CYBERARK_API_KEY" => Some(api_key.into()), + "CYBERARK_CLIENT_CERT" => Some(certificate.to_str().unwrap().into()), + "CYBERARK_CLIENT_KEY" => Some(key.to_str().unwrap().into()), + _ => None, + }), + true, + ) + .unwrap(); + + assert!(server.received_requests().await.unwrap().is_empty()); + assert_eq!( + manager + .async_read_secret("key") + .await + .unwrap() + .unwrap() + .expose(), + " value\n" + ); +} + +#[rstest] +#[case::certificate_only("", "client.crt")] +#[case::key_only("", "client.key")] +#[case::certificate_with_api_key("k3y", "client.crt")] +#[case::key_with_api_key("k3y", "client.key")] +fn invalid_client_identity_is_not_ignored( + client_identity_directory: tempfile::TempDir, + #[case] api_key: &'static str, + #[case] invalid_file: &str, +) { + std::fs::write( + client_identity_directory.path().join(invalid_file), + "not PEM", + ) + .unwrap(); + let certificate = client_identity_directory.path().join("client.crt"); + let key = client_identity_directory.path().join("client.key"); + + let result = CyberArkSecretManager::new( + Arc::new(move |name: &str| match name { + "CYBERARK_API_KEY" => Some(api_key.into()), + "CYBERARK_CLIENT_CERT" => Some(certificate.to_str().unwrap().into()), + "CYBERARK_CLIENT_KEY" => Some(key.to_str().unwrap().into()), + _ => None, + }), + true, + ); + + assert!(matches!(result, Err(Error::ClientCertificate))); +} + +#[rstest] +#[case::certificate_only("")] +#[case::certificate_and_api_key("k3y")] +fn client_identity_does_not_bypass_the_enterprise_requirement(#[case] api_key: &'static str) { + let result = CyberArkSecretManager::new( + Arc::new(move |name: &str| match name { + "CYBERARK_API_KEY" => Some(api_key.into()), + "CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()), + "CYBERARK_CLIENT_KEY" => Some("/missing/key".into()), + _ => None, + }), + false, + ); + + assert!(matches!(result, Err(Error::EnterpriseRequired))); +} + +#[rstest] +#[tokio::test] +async fn new_reads_environment_defaults_end_to_end() { + let server = MockServer::start().await; + Mock::given(path("/authn/default/admin/authenticate")) + .and(body_string("k3y")) + .respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON)) + .mount(&server) + .await; + Mock::given(path("/secrets/default/variable/key")) + .respond_with(ResponseTemplate::new(200).set_body_string("value")) + .mount(&server) + .await; + let endpoint = server.uri(); + let manager = CyberArkSecretManager::new( + Arc::new(move |name: &str| match name { + "CYBERARK_API_BASE" => Some(endpoint.clone()), + "CYBERARK_API_KEY" => Some("k3y".into()), + _ => None, + }), + true, + ) + .unwrap(); + assert_eq!( + manager + .async_read_secret("key") + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); +} + +#[rstest] +fn new_reports_missing_client_certificate_files() { + assert!(matches!( + CyberArkSecretManager::new( + Arc::new(|name: &str| match name { + "CYBERARK_API_KEY" => Some("k3y".into()), + "CYBERARK_CLIENT_CERT" => Some("/missing/cert".into()), + "CYBERARK_CLIENT_KEY" => Some("/missing/key".into()), + _ => None, + }), + true + ), + Err(Error::ClientCertificate) + )); +} + +#[rstest] +#[tokio::test] +async fn trailing_slash_endpoint_preserves_base_path() { + let server = MockServer::start().await; + Mock::given(path("/prefix/authn/acct/admin/authenticate")) + .and(body_string("k3y")) + .respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON)) + .expect(1) + .mount(&server) + .await; + Mock::given(path("/prefix/secrets/acct/variable/key")) + .respond_with(ResponseTemplate::new(200).set_body_string("value")) + .mount(&server) + .await; + let endpoint = format!("{}/prefix/", server.uri()).parse().unwrap(); + let manager = CyberArkSecretManager::with_client( + reqwest::Client::new(), + endpoint, + "acct".into(), + "admin".into(), + SecretValue::new("k3y"), + Some(Duration::from_secs(60)), + ); + assert_eq!( + manager + .async_read_secret("key") + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); +} diff --git a/litellm-rust/crates/secrets-cyberark/tests/secret_manager/reads.rs b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/reads.rs new file mode 100644 index 00000000000..3e5515841ed --- /dev/null +++ b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/reads.rs @@ -0,0 +1,145 @@ +use super::*; + +#[rstest] +#[tokio::test] +async fn a_rejected_refreshed_token_surfaces_the_error_without_another_retry() { + let server = MockServer::start().await; + mount_auth(&server, 2).await; + Mock::given(path("/secrets/acct/variable/first")) + .respond_with(ResponseTemplate::new(200).set_body_string("first-value")) + .expect(1) + .mount(&server) + .await; + Mock::given(path("/secrets/acct/variable/second")) + .respond_with(ResponseTemplate::new(401)) + .expect(2) + .mount(&server) + .await; + let manager = manager(&server, Duration::from_secs(60)); + assert_eq!( + manager + .async_read_secret("first") + .await + .unwrap() + .unwrap() + .expose(), + "first-value" + ); + + let result = tokio::time::timeout(Duration::from_secs(5), manager.async_read_secret("second")) + .await + .expect("authentication retries must terminate"); + + assert!(matches!(result, Err(Error::Status(401)))); +} + +#[rstest] +#[case::not_found(404)] +#[case::unauthorized(401)] +#[case::forbidden(403)] +#[case::server_error(500)] +#[tokio::test] +async fn failed_reads_are_not_cached(#[case] status: u16) { + let server = MockServer::start().await; + mount_auth(&server, 1).await; + let failing = Mock::given(path("/secrets/acct/variable/key")) + .respond_with(ResponseTemplate::new(status)) + .expect(1) + .mount_as_scoped(&server) + .await; + let manager = manager(&server, Duration::from_secs(60)); + let result = manager.async_read_secret("key").await; + if status == 404 { + assert_eq!(result.unwrap(), None); + } else { + assert!(matches!(result, Err(Error::Status(actual)) if actual == status)); + } + drop(failing); + Mock::given(path("/secrets/acct/variable/key")) + .respond_with(ResponseTemplate::new(200).set_body_string("recovered")) + .expect(1) + .mount(&server) + .await; + for _ in 0..2 { + assert_eq!( + manager + .async_read_secret("key") + .await + .unwrap() + .unwrap() + .expose(), + "recovered" + ); + } +} + +#[rstest] +#[tokio::test] +async fn expired_tokens_and_secrets_are_fetched_again() { + let server = MockServer::start().await; + mount_auth(&server, 2).await; + Mock::given(path("/secrets/acct/variable/key")) + .respond_with(ResponseTemplate::new(200).set_body_string("value")) + .expect(2) + .mount(&server) + .await; + let manager = manager(&server, Duration::from_millis(1)); + for _ in 0..2 { + assert!(manager.async_read_secret("key").await.unwrap().is_some()); + tokio::time::sleep(Duration::from_millis(5)).await; + } +} + +#[rstest] +#[case::plain("OPENAI_API_KEY")] +#[case::path("team/app/key")] +#[case::punctuation("a b+c.d-e_f~g")] +#[case::quote("needs \"quote\"")] +#[tokio::test] +async fn secret_names_use_python_quote_encoding(parity_fixture: ParityFixture, #[case] name: &str) { + let secret = parity_fixture + .secrets + .iter() + .find(|secret| secret.name == name) + .unwrap(); + let server = MockServer::start().await; + mount_auth(&server, 1).await; + Mock::given(RawPath(secret.path.clone())) + .respond_with(ResponseTemplate::new(200).set_body_string("value")) + .expect(1) + .mount(&server) + .await; + assert_eq!( + manager(&server, Duration::from_secs(60)) + .async_read_secret(name) + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); +} + +#[rstest] +#[case::parent("../etc")] +#[case::embedded_parent("team/../etc")] +#[case::control("key\n")] +#[tokio::test] +async fn unsafe_names_fail_before_http_calls(#[case] name: &str) { + let server = MockServer::start().await; + let manager = manager(&server, Duration::from_secs(60)); + assert!(matches!( + manager.async_read_secret(name).await, + Err(Error::Operation( + litellm_secrets_types::Error::UnsafeSecretName + )) + )); + assert!(matches!( + manager + .async_write_secret(name, &SecretValue::new("v"), None) + .await, + Err(Error::Operation( + litellm_secrets_types::Error::UnsafeSecretName + )) + )); +} diff --git a/litellm-rust/crates/secrets-cyberark/tests/secret_manager/support.rs b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/support.rs new file mode 100644 index 00000000000..f5ba7a63273 --- /dev/null +++ b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/support.rs @@ -0,0 +1,70 @@ +use super::*; + +pub(super) const TOKEN_JSON: &str = r#"{"protected":"p","payload":"q","signature":"s"}"#; + +#[derive(Deserialize)] +pub(super) struct ParityFixture { + pub(super) account: String, + pub(super) username: String, + pub(super) api_key: String, + pub(super) authenticate_path: String, + pub(super) token_json: String, + pub(super) authorization_header: String, + pub(super) policy_path: String, + pub(super) secrets: Vec, +} + +#[derive(Deserialize)] +pub(super) struct ParitySecret { + pub(super) name: String, + pub(super) path: String, + pub(super) policy_body: String, +} + +#[derive(Debug)] +pub(super) struct RawPath(pub(super) String); + +impl Match for RawPath { + fn matches(&self, request: &Request) -> bool { + request.url.path() == self.0 + } +} + +#[fixture] +pub(super) fn parity_fixture() -> ParityFixture { + serde_json::from_str(include_str!("../fixtures/parity.json")).unwrap() +} + +#[fixture] +pub(super) fn client_identity_directory() -> tempfile::TempDir { + let identity = rcgen::generate_simple_self_signed(vec!["localhost".into()]).unwrap(); + let directory = tempfile::tempdir().unwrap(); + std::fs::write(directory.path().join("client.crt"), identity.cert.pem()).unwrap(); + std::fs::write( + directory.path().join("client.key"), + identity.signing_key.serialize_pem(), + ) + .unwrap(); + directory +} + +pub(super) fn manager(server: &MockServer, ttl: Duration) -> CyberArkSecretManager { + CyberArkSecretManager::with_client( + reqwest::Client::new(), + server.uri().parse().unwrap(), + "acct".into(), + "admin".into(), + SecretValue::new("k3y"), + Some(ttl), + ) +} + +pub(super) async fn mount_auth(server: &MockServer, expected: u64) { + Mock::given(method("POST")) + .and(path("/authn/acct/admin/authenticate")) + .and(body_string("k3y")) + .respond_with(ResponseTemplate::new(200).set_body_string(TOKEN_JSON)) + .expect(expected) + .mount(server) + .await; +} diff --git a/litellm-rust/crates/secrets-cyberark/tests/secret_manager/writes.rs b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/writes.rs new file mode 100644 index 00000000000..331a26c6119 --- /dev/null +++ b/litellm-rust/crates/secrets-cyberark/tests/secret_manager/writes.rs @@ -0,0 +1,450 @@ +use super::*; + +#[rstest] +#[tokio::test] +async fn rejected_write_token_is_reauthenticated_once() { + let server = MockServer::start().await; + mount_auth(&server, 2).await; + Mock::given(path("/policies/acct/policy/root")) + .respond_with(ResponseTemplate::new(409)) + .expect(1) + .mount(&server) + .await; + let attempts = Arc::new(AtomicUsize::new(0)); + let attempts_for_response = Arc::clone(&attempts); + Mock::given(method("POST")) + .and(path("/secrets/acct/variable/key")) + .and(body_string("value")) + .respond_with(move |_: &Request| { + if attempts_for_response.fetch_add(1, Ordering::SeqCst) == 0 { + ResponseTemplate::new(401) + } else { + ResponseTemplate::new(200) + } + }) + .expect(2) + .mount(&server) + .await; + let manager = manager(&server, Duration::from_secs(600)); + + manager + .async_write_secret("key", &SecretValue::new("value"), None) + .await + .unwrap(); + assert_eq!(attempts.load(Ordering::SeqCst), 2); + assert_eq!( + manager + .async_read_secret("key") + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); +} + +#[rstest] +#[case::created(201)] +#[case::already_exists(409)] +#[case::unprocessable(422)] +#[case::server_error(500)] +#[tokio::test] +async fn writes_tolerate_policy_status_and_cache_value(#[case] policy_status: u16) { + let server = MockServer::start().await; + mount_auth(&server, 1).await; + Mock::given(path("/policies/acct/policy/root")) + .and(header("content-type", "application/x-yaml")) + .and(body_string("- !variable \"team/app\"\n")) + .respond_with(ResponseTemplate::new(policy_status)) + .expect(1) + .mount(&server) + .await; + Mock::given(path("/secrets/acct/variable/team%2Fapp")) + .and(body_string("v")) + .respond_with(ResponseTemplate::new(200)) + .expect(1) + .mount(&server) + .await; + let manager = manager(&server, Duration::from_secs(60)); + manager + .async_write_secret("team/app", &SecretValue::new("v"), None) + .await + .unwrap(); + assert_eq!( + manager + .async_read_secret("team/app") + .await + .unwrap() + .unwrap() + .expose(), + "v" + ); +} + +#[rstest] +#[tokio::test] +async fn failed_value_write_is_not_cached() { + let server = MockServer::start().await; + mount_auth(&server, 1).await; + Mock::given(path("/policies/acct/policy/root")) + .respond_with(ResponseTemplate::new(409)) + .mount(&server) + .await; + Mock::given(path("/secrets/acct/variable/key")) + .and(body_string("v")) + .respond_with(ResponseTemplate::new(403)) + .expect(1) + .mount(&server) + .await; + Mock::given(path("/secrets/acct/variable/key")) + .respond_with(ResponseTemplate::new(200).set_body_string("recovered")) + .expect(1) + .mount(&server) + .await; + let manager = manager(&server, Duration::from_secs(60)); + assert!(matches!( + manager + .async_write_secret("key", &SecretValue::new("v"), None) + .await, + Err(Error::Status(403)) + )); + assert_eq!( + manager + .async_read_secret("key") + .await + .unwrap() + .unwrap() + .expose(), + "recovered" + ); +} + +#[rstest] +#[tokio::test] +async fn delete_invalidates_cache_and_reports_not_supported() { + let server = MockServer::start().await; + mount_auth(&server, 1).await; + Mock::given(path("/secrets/acct/variable/key")) + .respond_with(ResponseTemplate::new(200).set_body_string("v")) + .expect(2) + .mount(&server) + .await; + let manager = manager(&server, Duration::from_secs(60)); + assert_eq!( + manager + .async_read_secret("key") + .await + .unwrap() + .unwrap() + .expose(), + "v" + ); + assert_eq!( + manager.async_delete_secret("key", Some(7)).await.unwrap(), + DeleteOutcome::NotSupported + ); + assert_eq!( + manager + .async_read_secret("key") + .await + .unwrap() + .unwrap() + .expose(), + "v" + ); +} + +#[rstest] +#[tokio::test] +async fn writes_match_python_parity_fixture(parity_fixture: ParityFixture) { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path(&parity_fixture.authenticate_path)) + .and(body_string(&parity_fixture.api_key)) + .respond_with(ResponseTemplate::new(200).set_body_string(&parity_fixture.token_json)) + .expect(1) + .mount(&server) + .await; + let manager = CyberArkSecretManager::with_client( + reqwest::Client::new(), + server.uri().parse().unwrap(), + parity_fixture.account, + parity_fixture.username, + SecretValue::new(parity_fixture.api_key), + Some(Duration::from_secs(60)), + ); + for secret in parity_fixture.secrets { + Mock::given(method("POST")) + .and(path(&parity_fixture.policy_path)) + .and(header( + "authorization", + &parity_fixture.authorization_header, + )) + .and(header("content-type", "application/x-yaml")) + .and(body_string(&secret.policy_body)) + .respond_with(ResponseTemplate::new(201)) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(RawPath(secret.path)) + .and(header( + "authorization", + &parity_fixture.authorization_header, + )) + .and(body_string("value")) + .respond_with(ResponseTemplate::new(201)) + .expect(1) + .mount(&server) + .await; + manager + .async_write_secret(&secret.name, &SecretValue::new("value"), None) + .await + .unwrap(); + assert_eq!( + manager + .async_read_secret(&secret.name) + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); + } +} + +#[rstest] +#[tokio::test] +#[ignore] +async fn live_conjur_round_trip() { + let endpoint: reqwest::Url = std::env::var("CYBERARK_API_BASE").unwrap().parse().unwrap(); + let account = std::env::var("CYBERARK_ACCOUNT").unwrap(); + let username = std::env::var("CYBERARK_USERNAME").unwrap(); + let api_key = SecretValue::new(std::env::var("CYBERARK_API_KEY").unwrap()); + let name = format!( + "{}-{}", + std::env::var("LITELLM_CONJUR_LIVE_SECRET_NAME").unwrap(), + std::time::SystemTime::now() + .duration_since(std::time::UNIX_EPOCH) + .unwrap() + .as_nanos() + ); + let manager = CyberArkSecretManager::with_client( + reqwest::Client::new(), + endpoint.clone(), + account.clone(), + username.clone(), + api_key.clone(), + Some(Duration::from_secs(60)), + ); + + assert!(manager.async_read_secret(&name).await.unwrap().is_none()); + for expected in ["first-Ï€\n", " second-Ï€ "] { + manager + .async_write_secret(&name, &SecretValue::new(expected), None) + .await + .unwrap(); + let verifier = CyberArkSecretManager::with_client( + reqwest::Client::new(), + endpoint.clone(), + account.clone(), + username.clone(), + api_key.clone(), + Some(Duration::from_secs(60)), + ); + assert_eq!( + verifier + .async_read_secret(&name) + .await + .unwrap() + .unwrap() + .expose(), + expected + ); + } + manager + .async_rotate_secret(&name, &name, &SecretValue::new("rotated-value")) + .await + .unwrap(); + let alias = format!("{name}-rotated"); + manager + .async_rotate_secret(&name, &alias, &SecretValue::new("new-alias-value")) + .await + .unwrap(); + assert_eq!( + manager + .async_read_secret(&name) + .await + .unwrap() + .unwrap() + .expose(), + "rotated-value" + ); + assert_eq!( + manager + .async_read_secret(&alias) + .await + .unwrap() + .unwrap() + .expose(), + "new-alias-value" + ); +} + +#[rstest] +#[case::same_alias("old")] +#[case::new_alias("new")] +#[tokio::test] +async fn rotation_stores_the_replacement_and_retains_other_aliases(#[case] new_name: &'static str) { + use std::sync::atomic::{AtomicBool, Ordering}; + let server = MockServer::start().await; + mount_auth(&server, 1).await; + let written = Arc::new(AtomicBool::new(false)); + let read_state = written.clone(); + Mock::given(method("GET")) + .respond_with(move |request: &wiremock::Request| { + let name = request.url.path().rsplit('/').next().unwrap(); + let value = if name == new_name && read_state.load(Ordering::SeqCst) { + "new-value" + } else { + "old-value" + }; + ResponseTemplate::new(200).set_body_string(value) + }) + .expect(if new_name == "old" { 2 } else { 3 }) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/policies/acct/policy/root")) + .and(body_string(format!("- !variable \"{new_name}\"\n"))) + .respond_with(ResponseTemplate::new(201)) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path(format!("/secrets/acct/variable/{new_name}"))) + .and(body_string("new-value")) + .respond_with(move |_: &wiremock::Request| { + written.store(true, Ordering::SeqCst); + ResponseTemplate::new(201) + }) + .expect(1) + .mount(&server) + .await; + let manager = manager(&server, Duration::from_secs(60)); + manager + .async_rotate_secret("old", new_name, &SecretValue::new("new-value")) + .await + .unwrap(); + assert_eq!( + manager + .async_read_secret(new_name) + .await + .unwrap() + .unwrap() + .expose(), + "new-value" + ); + assert_eq!( + manager + .async_read_secret("old") + .await + .unwrap() + .unwrap() + .expose(), + if new_name == "old" { + "new-value" + } else { + "old-value" + } + ); + assert!( + server + .received_requests() + .await + .unwrap() + .iter() + .all(|request| request.method != "DELETE") + ); +} + +#[tokio::test] +async fn rotation_verifies_the_remote_value_instead_of_the_write_cache() { + let server = MockServer::start().await; + mount_auth(&server, 1).await; + Mock::given(method("GET")) + .respond_with(ResponseTemplate::new(200).set_body_string("unchanged")) + .expect(2) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/policies/acct/policy/root")) + .respond_with(ResponseTemplate::new(201)) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/secrets/acct/variable/new")) + .respond_with(ResponseTemplate::new(201)) + .expect(1) + .mount(&server) + .await; + let result = manager(&server, Duration::from_secs(60)) + .async_rotate_secret("old", "new", &SecretValue::new("replacement")) + .await; + assert!(matches!( + result, + Err(litellm_secrets_types::RotationError::Verification { + source: Error::Operation(litellm_secrets_types::Error::NewSecretMismatch), + .. + }) + )); +} + +#[rstest] +#[case::colon("foo: bar")] +#[case::comment("foo # bar")] +#[case::plain("plain-alias")] +#[case::email("team/user@example.com")] +#[case::quote("needs \"quote\"")] +#[case::backslash("a\\b")] +#[tokio::test] +async fn policy_writes_preserve_yaml_metacharacters_as_one_variable(#[case] name: &'static str) { + let server = MockServer::start().await; + mount_auth(&server, 1).await; + Mock::given(method("POST")) + .and(path("/policies/acct/policy/root")) + .respond_with(move |request: &wiremock::Request| { + let body = std::str::from_utf8(&request.body).unwrap(); + let scalar = body + .strip_prefix("- !variable ") + .unwrap() + .strip_suffix('\n') + .unwrap(); + assert_eq!(serde_json::from_str::(scalar).unwrap(), name); + ResponseTemplate::new(201) + }) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(body_string("value")) + .respond_with(ResponseTemplate::new(201)) + .expect(1) + .mount(&server) + .await; + let manager = manager(&server, Duration::from_secs(60)); + manager + .async_write_secret(name, &SecretValue::new("value"), None) + .await + .unwrap(); + assert_eq!( + manager + .async_read_secret(name) + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); +} diff --git a/litellm-rust/crates/secrets-google/AGENTS.md b/litellm-rust/crates/secrets-google/AGENTS.md new file mode 100644 index 00000000000..2ad447f15b3 --- /dev/null +++ b/litellm-rust/crates/secrets-google/AGENTS.md @@ -0,0 +1 @@ +- https://docs.cloud.google.com/secret-manager/docs/reference/rest/v1/projects.secrets.versions/access diff --git a/litellm-rust/crates/secrets-google/Cargo.toml b/litellm-rust/crates/secrets-google/Cargo.toml index daecf20ff9e..805eb80740d 100644 --- a/litellm-rust/crates/secrets-google/Cargo.toml +++ b/litellm-rust/crates/secrets-google/Cargo.toml @@ -6,14 +6,16 @@ license.workspace = true repository.workspace = true [dependencies] +moka.workspace = true +tokio.workspace = true litellm-auth-gcp = { workspace = true, features = ["google-sdk"] } litellm-secrets-types.workspace = true litellm-auth-types.workspace = true litellm-core-utils.workspace = true base64.workspace = true +crc32c = "0.6.8" serde_json.workspace = true thiserror.workspace = true -moka.workspace = true veil.workspace = true google-cloud-kms-v1 = "1.14.0" google-cloud-gax = { version = "1.14.0", default-features = false } @@ -24,5 +26,4 @@ reqwest.workspace = true [dev-dependencies] google-cloud-auth.workspace = true rstest.workspace = true -tokio.workspace = true wiremock = "0.6.5" diff --git a/litellm-rust/crates/secrets-google/src/error.rs b/litellm-rust/crates/secrets-google/src/error.rs index a94cc8c2de6..159e4983ddc 100644 --- a/litellm-rust/crates/secrets-google/src/error.rs +++ b/litellm-rust/crates/secrets-google/src/error.rs @@ -1,5 +1,9 @@ #[derive(thiserror::Error, veil::Redact)] pub enum Error { + #[error(transparent)] + Operation(#[from] litellm_secrets_types::Error), + #[error("secret manager operation timed out")] + Timeout, #[error("Google KMS client configuration failed")] Client( #[from] @@ -34,6 +38,8 @@ pub enum Error { RefreshInterval, #[error("payload is not valid base64")] Base64(#[from] base64::DecodeError), + #[error("Google Secret Manager payload checksum mismatch")] + Checksum, #[error("decrypted value is not UTF-8")] Utf8, #[error("invalid Google Secret Manager endpoint")] diff --git a/litellm-rust/crates/secrets-google/src/kms.rs b/litellm-rust/crates/secrets-google/src/kms.rs index 3a247edaa35..78ebae008f7 100644 --- a/litellm-rust/crates/secrets-google/src/kms.rs +++ b/litellm-rust/crates/secrets-google/src/kms.rs @@ -36,12 +36,10 @@ impl GoogleKms { } pub fn validate_environment(environment: &dyn Lookup) -> Result<(), Error> { - for key in [GOOGLE_APPLICATION_CREDENTIALS, GOOGLE_KMS_RESOURCE_NAME] { - if environment.get(key).is_none() { - return Err(Error::MissingEnvironment(key)); - } - } - Ok(()) + environment + .get(GOOGLE_KMS_RESOURCE_NAME) + .map(|_| ()) + .ok_or(Error::MissingEnvironment(GOOGLE_KMS_RESOURCE_NAME)) } pub async fn load_google_kms( @@ -52,13 +50,16 @@ pub async fn load_google_kms( return Ok(None); } validate_environment(environment.as_ref())?; - let credentials = environment - .get(GOOGLE_APPLICATION_CREDENTIALS) - .ok_or(Error::MissingEnvironment(GOOGLE_APPLICATION_CREDENTIALS))?; let resource_name = environment .get(GOOGLE_KMS_RESOURCE_NAME) .ok_or(Error::MissingEnvironment(GOOGLE_KMS_RESOURCE_NAME))?; - let credentials = auth::credentials(None, Some(SecretValue::new(credentials)), environment); + let credentials = auth::credentials( + None, + environment + .get(GOOGLE_APPLICATION_CREDENTIALS) + .map(SecretValue::new), + environment, + ); let client = KeyManagementService::builder() .with_credentials(credentials) .build() diff --git a/litellm-rust/crates/secrets-google/src/secret_manager.rs b/litellm-rust/crates/secrets-google/src/secret_manager.rs index 3c34d9cbcc4..b8787999e12 100644 --- a/litellm-rust/crates/secrets-google/src/secret_manager.rs +++ b/litellm-rust/crates/secrets-google/src/secret_manager.rs @@ -2,8 +2,9 @@ use std::{sync::Arc, time::Duration}; use base64::{Engine, engine::general_purpose::STANDARD}; use litellm_core_utils::settings::Lookup; -use litellm_secrets_types::{Secret, SecretValue}; -use moka::future::Cache; +use litellm_secrets_types::{ + BaseSecretManager, GoogleOperationContext, Secret, SecretCache, SecretValue, +}; use serde::Deserialize; use litellm_auth_gcp::GoogleCredentials; @@ -26,7 +27,8 @@ pub struct GoogleSecretManager { credentials: Arc, endpoint: reqwest::Url, project: String, - cache: Cache, + cache: SecretCache, + python_misses: moka::future::Cache, always_read: bool, } @@ -38,6 +40,8 @@ struct Response { #[derive(Deserialize)] struct Payload { data: Option, + #[serde(rename = "dataCrc32c")] + data_crc32c: Option, } impl GoogleSecretManager { @@ -59,16 +63,17 @@ impl GoogleSecretManager { let ttl = refresh_interval .filter(|ttl| !ttl.is_zero()) .unwrap_or(DEFAULT_CACHE_TTL); - let cache = Cache::builder() - .max_capacity(CACHE_CAPACITY) - .time_to_live(ttl) - .build(); + let cache = SecretCache::new(CACHE_CAPACITY, ttl); Ok(Self { client, credentials: Arc::new(credentials), endpoint, project, cache, + python_misses: moka::future::Cache::builder() + .max_capacity(CACHE_CAPACITY) + .time_to_live(ttl) + .build(), always_read, }) } @@ -116,11 +121,38 @@ impl GoogleSecretManager { &self, name: &str, ) -> Result, Error> { - if !self.always_read - && let Some(cached) = self.cache.get(name).await - { - return Ok(Some(Secret::String(cached))); + BaseSecretManager::async_read_secret(self, name, &GoogleOperationContext::default()) + .await + .map(|value| value.map(Secret::String)) + } + + pub async fn get_secret_for_python(&self, name: &str) -> Result, Error> { + if !self.always_read && self.python_misses.get(name).await.is_some() { + return Ok(None); } + let result = self.get_secret_from_google_secret_manager(name).await; + if matches!( + result, + Ok(None) | Err(Error::Status(_) | Error::MissingPayload) + ) { + self.python_misses.insert(name.to_owned(), ()).await; + } + match result { + Ok(None) => Err(Error::Status(404)), + result => result, + } + } + + async fn read(&self, name: &str) -> Result, Error> { + if self.always_read { + return self.read_uncached(name).await; + } + self.cache + .read(name.to_owned(), self.read_uncached(name)) + .await + } + + async fn read_uncached(&self, name: &str) -> Result, Error> { let url = self .endpoint .join(&format!( @@ -145,13 +177,39 @@ impl GoogleSecretManager { return Err(Error::Status(response.status().as_u16())); } let response: Response = response.json().await?; - let Some(data) = response.payload.and_then(|payload| payload.data) else { + let Some(payload) = response.payload else { + return Err(Error::MissingPayload); + }; + let Some(data) = payload.data else { return Err(Error::MissingPayload); }; let bytes = STANDARD.decode(data)?; + if let Some(expected) = payload.data_crc32c { + let expected = expected.parse::().map_err(|_| Error::Checksum)?; + if crc32c::crc32c(&bytes) != expected { + return Err(Error::Checksum); + } + } let plaintext = String::from_utf8(bytes).map_err(|_| Error::Utf8)?; let value = SecretValue::new(plaintext); - self.cache.insert(name.to_owned(), value.clone()).await; - Ok(Some(Secret::String(value))) + Ok(Some(value)) + } +} + +impl BaseSecretManager for GoogleSecretManager { + type Error = Error; + type Context = GoogleOperationContext; + + async fn async_read_secret( + &self, + name: &str, + context: &Self::Context, + ) -> Result, Error> { + match context.timeout { + Some(timeout) => tokio::time::timeout(timeout, self.read(name)) + .await + .map_err(|_| Error::Timeout)?, + None => self.read(name).await, + } } } diff --git a/litellm-rust/crates/secrets-google/tests/kms.rs b/litellm-rust/crates/secrets-google/tests/kms.rs index 9739a46b10c..0ccdcc9036f 100644 --- a/litellm-rust/crates/secrets-google/tests/kms.rs +++ b/litellm-rust/crates/secrets-google/tests/kms.rs @@ -42,22 +42,41 @@ async fn google_kms_decrypts_using_the_configured_resource() { #[case::unset(None)] #[case::disabled(Some(false))] #[tokio::test] -async fn disabled_google_kms_loader_does_not_require_environment_configuration( +async fn disabled_google_kms_loader_does_not_read_environment_configuration( #[case] enabled: Option, ) { use std::sync::Arc; assert!( - litellm_secrets_google::load_google_kms(enabled, Arc::new(|_: &str| None)) - .await - .unwrap() - .is_none() + litellm_secrets_google::load_google_kms( + enabled, + Arc::new(|name: &str| panic!("disabled Google KMS read {name}")), + ) + .await + .unwrap() + .is_none() ); } #[rstest] -#[case::credentials_missing(None, None, "GOOGLE_APPLICATION_CREDENTIALS")] -#[case::resource_missing(Some("credentials"), None, "GOOGLE_KMS_RESOURCE_NAME")] -fn enabled_google_kms_requires_all_environment_values( +#[tokio::test] +async fn enabled_google_kms_loader_accepts_application_default_credentials() { + use std::sync::Arc; + let environment = Arc::new(|name: &str| { + (name == "GOOGLE_KMS_RESOURCE_NAME") + .then(|| "projects/project/locations/global/keyRings/ring/cryptoKeys/key".to_owned()) + }); + + assert!( + litellm_secrets_google::load_google_kms(Some(true), environment) + .await + .unwrap() + .is_some() + ); +} + +#[rstest] +#[case::resource_missing(None, None, "GOOGLE_KMS_RESOURCE_NAME")] +fn enabled_google_kms_requires_resource_name( #[case] credentials: Option<&str>, #[case] resource: Option<&str>, #[case] missing: &'static str, @@ -75,9 +94,13 @@ fn enabled_google_kms_requires_all_environment_values( } #[rstest] -fn complete_google_kms_environment_is_valid() { +#[case::service_account_file(Some("credentials"))] +#[case::application_default_credentials(None)] +fn google_kms_environment_is_valid_without_required_credential_file( + #[case] credentials: Option<&str>, +) { let environment = |name: &str| match name { - "GOOGLE_APPLICATION_CREDENTIALS" => Some("credentials".to_owned()), + "GOOGLE_APPLICATION_CREDENTIALS" => credentials.map(str::to_owned), "GOOGLE_KMS_RESOURCE_NAME" => Some("resource".to_owned()), _ => None, }; diff --git a/litellm-rust/crates/secrets-google/tests/secret_manager.rs b/litellm-rust/crates/secrets-google/tests/secret_manager.rs index 867dcf935c1..0d7efc4b1b3 100644 --- a/litellm-rust/crates/secrets-google/tests/secret_manager.rs +++ b/litellm-rust/crates/secrets-google/tests/secret_manager.rs @@ -61,12 +61,43 @@ async fn successful_reads_use_auth_latest_version_and_cache_including_empty_valu } } +#[tokio::test] +async fn matching_checksum_is_accepted_and_cached() { + let server = MockServer::start().await; + let value = "private-value"; + Mock::given(path( + "/v1/projects/project/secrets/key/versions/latest:access", + )) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "payload": { + "data": STANDARD.encode(value), + "dataCrc32c": crc32c::crc32c(value.as_bytes()).to_string() + } + }))) + .expect(1) + .mount(&server) + .await; + let manager = manager(&server, false, Duration::from_secs(60)); + for _ in 0..2 { + assert_eq!( + manager + .get_secret_from_google_secret_manager("key") + .await + .unwrap() + .unwrap() + .as_str(), + Some(value) + ); + } +} + enum ExpectedReadFailure { Missing, Status(u16), MissingPayload, Base64, Utf8, + Checksum, } #[rstest] @@ -90,6 +121,11 @@ enum ExpectedReadFailure { serde_json::json!({"payload":{"data":STANDARD.encode([0xff])}}), ExpectedReadFailure::Utf8 )] +#[case::checksum_mismatch( + 200, + serde_json::json!({"payload":{"data":STANDARD.encode("corrupt"),"dataCrc32c":"0"}}), + ExpectedReadFailure::Checksum +)] #[tokio::test] async fn failed_or_missing_reads_are_not_cached( default_ttl: Duration, @@ -117,6 +153,7 @@ async fn failed_or_missing_reads_are_not_cached( } ExpectedReadFailure::Base64 => assert!(matches!(result, Err(Error::Base64(_)))), ExpectedReadFailure::Utf8 => assert!(matches!(result, Err(Error::Utf8))), + ExpectedReadFailure::Checksum => assert!(matches!(result, Err(Error::Checksum))), } drop(failing); Mock::given(path( @@ -235,3 +272,129 @@ async fn cache_preserves_raw_values(default_ttl: Duration, #[case] raw: &str) { ); } } + +#[tokio::test] +async fn trait_read_limits_the_operation_duration() { + use litellm_secrets_types::{BaseSecretManager, GoogleOperationContext}; + use std::time::Duration; + let server = MockServer::start().await; + Mock::given(wiremock::matchers::method("GET")) + .respond_with(ResponseTemplate::new(200).set_delay(Duration::from_secs(1))) + .mount(&server) + .await; + let manager = manager(&server, false, Duration::from_secs(60)); + let context = GoogleOperationContext { + timeout: Some(Duration::from_millis(30)), + }; + assert!(matches!( + BaseSecretManager::async_read_secret(&manager, "key", &context).await, + Err(Error::Timeout) + )); +} + +#[tokio::test] +async fn concurrent_reads_share_one_secret_request() { + let server = MockServer::start().await; + Mock::given(wiremock::matchers::method("GET")) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(serde_json::json!({"payload": {"data": STANDARD.encode("value")}})) + .set_delay(Duration::from_millis(20)), + ) + .expect(1) + .mount(&server) + .await; + let manager = manager(&server, false, Duration::from_secs(60)); + let (first, second) = tokio::join!( + manager.get_secret_from_google_secret_manager("key"), + manager.get_secret_from_google_secret_manager("key") + ); + assert_eq!(first.unwrap().unwrap().as_str(), Some("value")); + assert_eq!(second.unwrap().unwrap().as_str(), Some("value")); +} + +#[rstest] +#[case::missing(404, serde_json::json!({}))] +#[case::failure(403, serde_json::json!({}))] +#[case::no_payload(200, serde_json::json!({"payload":{}}))] +#[tokio::test] +async fn python_reads_reuse_cached_absence_until_expiry( + #[case] status: u16, + #[case] body: serde_json::Value, + #[values(false, true)] always_read: bool, +) { + let server = MockServer::start().await; + let manager = manager(&server, always_read, Duration::from_secs(60)); + let failing = Mock::given(path( + "/v1/projects/project/secrets/key/versions/latest:access", + )) + .respond_with(ResponseTemplate::new(status).set_body_json(body)) + .expect(1) + .mount_as_scoped(&server) + .await; + let result = manager.get_secret_for_python("key").await; + match status { + 404 => assert!(matches!(result, Err(Error::Status(404)))), + 403 => assert!(matches!(result, Err(Error::Status(403)))), + 200 => assert!(matches!(result, Err(Error::MissingPayload))), + _ => unreachable!(), + } + drop(failing); + Mock::given(path( + "/v1/projects/project/secrets/key/versions/latest:access", + )) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(serde_json::json!({"payload":{"data":STANDARD.encode("recovered")}})), + ) + .expect(u64::from(always_read)) + .mount(&server) + .await; + assert_eq!( + manager + .get_secret_for_python("key") + .await + .unwrap() + .as_ref() + .and_then(|value| value.as_str()), + always_read.then_some("recovered") + ); +} + +#[tokio::test] +async fn python_cached_absence_expires_and_allows_recovery() { + let server = MockServer::start().await; + let manager = manager(&server, false, Duration::from_millis(20)); + let missing = Mock::given(path( + "/v1/projects/project/secrets/key/versions/latest:access", + )) + .respond_with(ResponseTemplate::new(404)) + .expect(1) + .mount_as_scoped(&server) + .await; + assert!(matches!( + manager.get_secret_for_python("key").await, + Err(Error::Status(404)) + )); + drop(missing); + tokio::time::sleep(Duration::from_millis(40)).await; + Mock::given(path( + "/v1/projects/project/secrets/key/versions/latest:access", + )) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(serde_json::json!({"payload":{"data":STANDARD.encode("recovered")}})), + ) + .expect(1) + .mount(&server) + .await; + assert_eq!( + manager + .get_secret_for_python("key") + .await + .unwrap() + .unwrap() + .as_str(), + Some("recovered") + ); +} diff --git a/litellm-rust/crates/secrets-hashicorp/AGENTS.md b/litellm-rust/crates/secrets-hashicorp/AGENTS.md new file mode 100644 index 00000000000..ba1d5f6e330 --- /dev/null +++ b/litellm-rust/crates/secrets-hashicorp/AGENTS.md @@ -0,0 +1 @@ +- https://developer.hashicorp.com/vault/api-docs/secret/kv/kv-v2 diff --git a/litellm-rust/crates/secrets-hashicorp/Cargo.toml b/litellm-rust/crates/secrets-hashicorp/Cargo.toml index c646e02ef09..c049ba127e5 100644 --- a/litellm-rust/crates/secrets-hashicorp/Cargo.toml +++ b/litellm-rust/crates/secrets-hashicorp/Cargo.toml @@ -8,11 +8,10 @@ repository.workspace = true [dependencies] litellm-core-utils.workspace = true litellm-secrets-types.workspace = true -moka.workspace = true rustify.workspace = true rustify_derive.workspace = true serde.workspace = true -serde_json.workspace = true +serde_json = { workspace = true, features = ["raw_value"] } thiserror.workspace = true tokio.workspace = true vaultrs.workspace = true diff --git a/litellm-rust/crates/secrets-hashicorp/src/error.rs b/litellm-rust/crates/secrets-hashicorp/src/error.rs index 664f9ab18dd..1813cb486bc 100644 --- a/litellm-rust/crates/secrets-hashicorp/src/error.rs +++ b/litellm-rust/crates/secrets-hashicorp/src/error.rs @@ -3,9 +3,9 @@ pub enum Error { #[error("HashiCorp Vault requires an enterprise license")] EnterpriseRequired, #[error("invalid secret name")] - InvalidSecretName(#[from] litellm_secrets_types::Error), - #[error("HashiCorp Vault received an incompatible operation context")] - InvalidOperationContext, + InvalidSecretName(litellm_secrets_types::Error), + #[error(transparent)] + Operation(#[from] litellm_secrets_types::Error), #[error("HashiCorp Vault client failed")] Client( #[from] @@ -31,6 +31,10 @@ pub enum Error { MalformedPayload, #[error("HashiCorp Vault secret value is not a string")] NonStringValue, + #[error("HashiCorp Vault data key conflicts with description")] + DataKeyConflictsWithDescription, + #[error("HashiCorp Vault secret version exceeds CAS range")] + CasVersionOverflow, #[error("HashiCorp Vault operation timed out")] Timeout, #[error("invalid HashiCorp Vault refresh interval")] diff --git a/litellm-rust/crates/secrets-hashicorp/src/lib.rs b/litellm-rust/crates/secrets-hashicorp/src/lib.rs index 0c2b05647f8..7aab5a38b01 100644 --- a/litellm-rust/crates/secrets-hashicorp/src/lib.rs +++ b/litellm-rust/crates/secrets-hashicorp/src/lib.rs @@ -7,4 +7,4 @@ pub mod secret_manager; pub use config::{AppRoleAuth, HashicorpVaultConfig, TlsCertAuth}; pub use error::Error; -pub use secret_manager::{HashicorpVault, SecretLocation}; +pub use secret_manager::{HashicorpVault, RawOperationError, SecretLocation}; diff --git a/litellm-rust/crates/secrets-hashicorp/src/secret_manager.rs b/litellm-rust/crates/secrets-hashicorp/src/secret_manager.rs index a4d1efb4ff6..f0d4fe8c0b4 100644 --- a/litellm-rust/crates/secrets-hashicorp/src/secret_manager.rs +++ b/litellm-rust/crates/secrets-hashicorp/src/secret_manager.rs @@ -1,3 +1,8 @@ +mod client; +mod raw; +mod read; +mod write; + use std::{ collections::HashMap, fmt, @@ -8,15 +13,16 @@ use std::{ use litellm_core_utils::settings::Lookup; use litellm_secrets_types::{ - BaseSecretManager, HashicorpOperationContext, SecretOperationContext, SecretValue, - SecretWriteContext, async_rotate_secret, validate_secret_name, + BaseSecretManager, HashicorpOperationContext, RotationError, SecretCache, SecretDeleter, + SecretRotator, SecretValue, SecretWriteContext, SecretWriter, async_rotate_secret, + validate_secret_name, }; -use moka::future::Cache; use rustify::errors::ClientError as RustifyClientError; use serde_json::Value; use tokio::sync::Mutex; use vaultrs::{ api, + api::kv2::requests::SetSecretRequestOptions, auth::approle, client::{Identity, VaultClient, VaultClientSettingsBuilder}, error::ClientError, @@ -25,6 +31,8 @@ use vaultrs::{ use crate::{Error, HashicorpVaultConfig, TlsCertAuth, cert_login::CertLoginRequest}; +pub use raw::RawOperationError; + const CACHE_CAPACITY: u64 = 200; #[derive(Clone)] @@ -49,7 +57,7 @@ struct CacheKey { #[derive(Clone)] pub struct HashicorpVault { config: HashicorpVaultConfig, - cache: Cache, + cache: SecretCache, auth_client: Arc>>, } @@ -79,10 +87,7 @@ impl HashicorpVault { if !enterprise_enabled { return Err(Error::EnterpriseRequired); } - let cache: Cache = Cache::builder() - .max_capacity(CACHE_CAPACITY) - .time_to_live(config.refresh_interval) - .build(); + let cache = SecretCache::new(CACHE_CAPACITY, config.refresh_interval); Ok(Self { config, cache, @@ -91,21 +96,21 @@ impl HashicorpVault { } pub fn secret_location(&self, secret_name: &str) -> Result { - self.secret_location_with_context(secret_name, &SecretOperationContext::default()) + self.secret_location_with_context(secret_name, &HashicorpOperationContext::default()) } pub fn secret_location_with_context( &self, secret_name: &str, - context: &SecretOperationContext, + context: &HashicorpOperationContext, ) -> Result { validate_secret_name(secret_name).map_err(Error::InvalidSecretName)?; - let operation: Option<&HashicorpOperationContext> = hashicorp_context(context)?; let path: String = [ - operation - .and_then(|operation| operation.path_prefix.as_deref()) - .and_then(path_component) - .or_else(|| self.config.path_prefix.clone()), + context + .path_prefix + .as_deref() + .or(self.config.path_prefix.as_deref()) + .and_then(path_component), Some(secret_name.to_owned()), ] .into_iter() @@ -113,11 +118,17 @@ impl HashicorpVault { .collect::>() .join("/"); Ok(SecretLocation { - namespace: self.config.secret_namespace().map(str::to_owned), - mount: operation - .and_then(|operation| operation.mount.as_deref()) + namespace: context + .namespace + .as_deref() + .or(self.config.secret_namespace()) + .and_then(path_component), + mount: context + .mount + .as_deref() + .or(Some(self.config.mount.as_str())) .and_then(path_component) - .unwrap_or_else(|| self.config.mount.clone()), + .unwrap_or_else(|| "secret".to_owned()), path, }) } @@ -125,249 +136,6 @@ impl HashicorpVault { pub fn config(&self) -> &HashicorpVaultConfig { &self.config } - - pub async fn async_read_secret(&self, secret_name: &str) -> Result, Error> { - self.async_read_secret_with_context(secret_name, &SecretOperationContext::default()) - .await - } - - pub async fn async_read_secret_with_context( - &self, - secret_name: &str, - context: &SecretOperationContext, - ) -> Result, Error> { - let location: SecretLocation = self.secret_location_with_context(secret_name, context)?; - let data_key: String = data_key(context)?; - let cache_key = CacheKey { - location: location.clone(), - data_key: data_key.clone(), - }; - if let Some(value) = self.cache.get(&cache_key).await { - return Ok(Some(value)); - } - let data: Option> = with_timeout(context, async { - let client: Arc = self.vault_client().await?; - match kv2::read(client.as_ref(), &location.mount, &location.path).await { - Ok(data) => Ok(Some(data)), - Err(error) if api_status(&error) == Some(404) => Ok(None), - Err(error) => Err(map_api_error(error, ErrorContext::Read)), - } - }) - .await?; - let Some(data) = data else { - return Ok(None); - }; - let Some(value) = data.get(&data_key) else { - return Ok(None); - }; - let value: &str = value.as_str().ok_or(Error::NonStringValue)?; - let value: SecretValue = SecretValue::new(value); - self.cache.insert(cache_key, value.clone()).await; - Ok(Some(value)) - } - - pub async fn async_write_secret( - &self, - secret_name: &str, - value: SecretValue, - description: Option<&str>, - ) -> Result { - self.async_write_secret_with_context( - secret_name, - &value, - &SecretWriteContext { - description: description.map(str::to_owned), - ..SecretWriteContext::default() - }, - ) - .await - } - - pub async fn async_write_secret_with_context( - &self, - secret_name: &str, - value: &SecretValue, - context: &SecretWriteContext, - ) -> Result { - let location: SecretLocation = - self.secret_location_with_context(secret_name, &context.operation)?; - let data_key: String = data_key(&context.operation)?; - let data: HashMap = match context.description.as_deref() { - Some(description) => [ - (data_key, Value::String(value.expose().to_owned())), - ( - "description".to_owned(), - Value::String(description.to_owned()), - ), - ] - .into_iter() - .collect(), - None => [(data_key, Value::String(value.expose().to_owned()))] - .into_iter() - .collect(), - }; - let metadata = with_timeout(&context.operation, async { - let client: Arc = self.vault_client().await?; - kv2::set(client.as_ref(), &location.mount, &location.path, &data) - .await - .map_err(|error| map_api_error(error, ErrorContext::Secret)) - }) - .await?; - self.cache.invalidate_all(); - serde_json::to_value(metadata) - .map_err(|source| Error::Client(ClientError::JsonParseError { source })) - } - - pub async fn async_delete_secret(&self, secret_name: &str) -> Result<(), Error> { - self.async_delete_secret_with_context(secret_name, &SecretOperationContext::default()) - .await - } - - pub async fn async_delete_secret_with_context( - &self, - secret_name: &str, - context: &SecretOperationContext, - ) -> Result<(), Error> { - let location: SecretLocation = self.secret_location_with_context(secret_name, context)?; - with_timeout(context, async { - let client: Arc = self.vault_client().await?; - kv2::delete_latest(client.as_ref(), &location.mount, &location.path) - .await - .map_err(|error| map_api_error(error, ErrorContext::Secret)) - }) - .await?; - self.cache.invalidate_all(); - Ok(()) - } - - pub async fn async_rotate_secret( - &self, - current_name: &str, - new_name: &str, - value: &SecretValue, - ) -> Result { - self.async_rotate_secret_with_context( - current_name, - new_name, - value, - &SecretOperationContext::default(), - ) - .await - } - - pub async fn async_rotate_secret_with_context( - &self, - current_name: &str, - new_name: &str, - value: &SecretValue, - context: &SecretOperationContext, - ) -> Result { - async_rotate_secret(self, current_name, new_name, value, context).await - } - - async fn vault_client(&self) -> Result, Error> { - let mut cached = self.auth_client.lock().await; - if let Some(entry) = cached.as_ref() - && entry - .expires_at - .is_none_or(|expires_at| expires_at > Instant::now()) - { - return Ok(entry.client.clone()); - } - - let (client, expires_at): (VaultClient, Option) = - match (self.config.approle.as_ref(), self.config.tls_cert.as_ref()) { - (Some(approle), _) => { - let login_client: VaultClient = - self.build_client(self.config.login_namespace(), "")?; - let auth = approle::login( - &login_client, - &approle.mount_path, - &approle.role_id, - approle.secret_id.expose(), - ) - .await - .map_err(|error| map_api_error(error, ErrorContext::Login))?; - ( - self.build_client(self.config.secret_namespace(), &auth.client_token)?, - token_expiry(auth.lease_duration), - ) - } - (None, Some(tls)) => { - let login_client: VaultClient = - self.build_client(self.config.login_namespace(), "")?; - let endpoint: CertLoginRequest = CertLoginRequest::new(tls.role.as_deref()); - let auth = api::auth(&login_client, endpoint) - .await - .map_err(|error| map_api_error(error, ErrorContext::Login))?; - ( - self.build_client(self.config.secret_namespace(), &auth.client_token)?, - token_expiry(auth.lease_duration), - ) - } - (None, None) => { - let token: SecretValue = - self.config.token.clone().ok_or(Error::NoAuthConfigured)?; - ( - self.build_client(self.config.secret_namespace(), token.expose())?, - None, - ) - } - }; - let client: Arc = Arc::new(client); - *cached = Some(CachedClient { - client: client.clone(), - expires_at, - }); - Ok(client) - } - - fn build_client(&self, namespace: Option<&str>, token: &str) -> Result { - let settings = VaultClientSettingsBuilder::default() - .address(&self.config.address) - .token(token.to_owned()) - .namespace(namespace.map(str::to_owned)) - .identity(identity_for(self.config.tls_cert.as_ref())?) - .ca_certs(Vec::new()) - .verify(true) - .build() - .map_err(|message| Error::ClientSettings { - message: message.to_string(), - })?; - VaultClient::new(settings).map_err(Error::Client) - } -} - -impl BaseSecretManager for HashicorpVault { - type Error = Error; - type WriteResponse = Value; - type DeleteResponse = (); - - async fn async_read_secret( - &self, - name: &str, - context: &SecretOperationContext, - ) -> Result, Error> { - HashicorpVault::async_read_secret_with_context(self, name, context).await - } - - async fn async_write_secret( - &self, - name: &str, - value: &SecretValue, - context: &SecretWriteContext, - ) -> Result { - HashicorpVault::async_write_secret_with_context(self, name, value, context).await - } - - async fn async_delete_secret( - &self, - name: &str, - _recovery_window_in_days: Option, - context: &SecretOperationContext, - ) -> Result<(), Error> { - HashicorpVault::async_delete_secret_with_context(self, name, context).await - } } #[derive(Clone, Copy)] @@ -377,37 +145,26 @@ enum ErrorContext { Secret, } -fn hashicorp_context( - context: &SecretOperationContext, -) -> Result, Error> { - match context { - SecretOperationContext::Hashicorp(context) => Ok(Some(context)), - SecretOperationContext::Default => Ok(None), - SecretOperationContext::Aws(_) | SecretOperationContext::Cyberark(_) => { - Err(Error::InvalidOperationContext) - } - } -} - fn path_component(value: &str) -> Option { let value: &str = value.trim().trim_matches('/'); (!value.is_empty()).then(|| value.to_owned()) } -fn data_key(context: &SecretOperationContext) -> Result { - Ok(hashicorp_context(context)? - .and_then(|context| context.data_key.as_deref()) +fn data_key(context: &HashicorpOperationContext) -> String { + context + .data_key + .as_deref() .map(str::trim) .filter(|data_key| !data_key.is_empty()) .map(str::to_owned) - .unwrap_or_else(|| "key".to_owned())) + .unwrap_or_else(|| "key".to_owned()) } async fn with_timeout( - context: &SecretOperationContext, + context: &HashicorpOperationContext, operation: impl Future>, ) -> Result { - match context.timeout() { + match context.timeout { Some(timeout) => tokio::time::timeout(timeout, operation) .await .map_err(|_| Error::Timeout)?, @@ -415,26 +172,6 @@ async fn with_timeout( } } -fn identity_for(tls: Option<&TlsCertAuth>) -> Result, Error> { - tls.map(|tls| { - let cert: Vec = std::fs::read(&tls.cert_path).map_err(|source| Error::TlsIdentity { - path: tls.cert_path.clone(), - message: source.to_string(), - })?; - let key: Vec = std::fs::read(&tls.key_path).map_err(|source| Error::TlsIdentity { - path: tls.key_path.clone(), - message: source.to_string(), - })?; - Identity::from_pem(&[cert.as_slice(), key.as_slice()].concat()).map_err(|source| { - Error::TlsIdentity { - path: tls.cert_path.clone(), - message: source.to_string(), - } - }) - }) - .transpose() -} - fn map_api_error(error: ClientError, context: ErrorContext) -> Error { match error { ClientError::APIError { code, .. } => match context { @@ -478,7 +215,3 @@ fn malformed_response(context: ErrorContext) -> Error { ErrorContext::Secret => Error::Client(ClientError::ResponseDataEmptyError), } } - -fn token_expiry(lease_duration: u64) -> Option { - (lease_duration > 0).then(|| Instant::now() + Duration::from_secs(lease_duration)) -} diff --git a/litellm-rust/crates/secrets-hashicorp/src/secret_manager/client.rs b/litellm-rust/crates/secrets-hashicorp/src/secret_manager/client.rs new file mode 100644 index 00000000000..c1dc073dcf8 --- /dev/null +++ b/litellm-rust/crates/secrets-hashicorp/src/secret_manager/client.rs @@ -0,0 +1,115 @@ +use super::*; + +impl HashicorpVault { + pub(super) async fn client_for_location( + &self, + location: &SecretLocation, + ) -> Result, Error> { + let client = self.vault_client().await?; + if client.settings.namespace == location.namespace { + return Ok(client); + } + self.build_client(location.namespace.as_deref(), &client.settings.token) + .map(Arc::new) + } + + pub(super) async fn vault_client(&self) -> Result, Error> { + let mut cached = self.auth_client.lock().await; + if let Some(entry) = cached.as_ref() + && entry + .expires_at + .is_none_or(|expires_at| expires_at > Instant::now()) + { + return Ok(entry.client.clone()); + } + + let (client, expires_at): (VaultClient, Option) = + match (self.config.approle.as_ref(), self.config.tls_cert.as_ref()) { + (Some(approle), _) => { + let login_client: VaultClient = + self.build_client(self.config.login_namespace(), "")?; + let auth = approle::login( + &login_client, + &approle.mount_path, + &approle.role_id, + approle.secret_id.expose(), + ) + .await + .map_err(|error| map_api_error(error, ErrorContext::Login))?; + ( + self.build_client(self.config.secret_namespace(), &auth.client_token)?, + token_expiry(auth.lease_duration), + ) + } + (None, Some(tls)) => { + let login_client: VaultClient = + self.build_client(self.config.login_namespace(), "")?; + let endpoint: CertLoginRequest = CertLoginRequest::new(tls.role.as_deref()); + let auth = api::auth(&login_client, endpoint) + .await + .map_err(|error| map_api_error(error, ErrorContext::Login))?; + ( + self.build_client(self.config.secret_namespace(), &auth.client_token)?, + token_expiry(auth.lease_duration), + ) + } + (None, None) => { + let token: SecretValue = + self.config.token.clone().ok_or(Error::NoAuthConfigured)?; + ( + self.build_client(self.config.secret_namespace(), token.expose())?, + None, + ) + } + }; + let client: Arc = Arc::new(client); + *cached = Some(CachedClient { + client: client.clone(), + expires_at, + }); + Ok(client) + } + + pub(super) fn build_client( + &self, + namespace: Option<&str>, + token: &str, + ) -> Result { + let settings = VaultClientSettingsBuilder::default() + .address(&self.config.address) + .token(token.to_owned()) + .namespace(namespace.map(str::to_owned)) + .identity(identity_for(self.config.tls_cert.as_ref())?) + .ca_certs(Vec::new()) + .verify(true) + .build() + .map_err(|message| Error::ClientSettings { + message: message.to_string(), + })?; + VaultClient::new(settings).map_err(Error::Client) + } +} + +fn identity_for(tls: Option<&TlsCertAuth>) -> Result, Error> { + tls.map(|tls| { + let cert: Vec = std::fs::read(&tls.cert_path).map_err(|source| Error::TlsIdentity { + path: tls.cert_path.clone(), + message: source.to_string(), + })?; + let key: Vec = std::fs::read(&tls.key_path).map_err(|source| Error::TlsIdentity { + path: tls.key_path.clone(), + message: source.to_string(), + })?; + Identity::from_pem(&[cert.as_slice(), key.as_slice()].concat()).map_err(|source| { + Error::TlsIdentity { + path: tls.cert_path.clone(), + message: source.to_string(), + } + }) + }) + .transpose() +} + +fn token_expiry(lease_duration: u64) -> Option { + (lease_duration > 0).then(|| Instant::now() + Duration::from_secs(lease_duration)) +} diff --git a/litellm-rust/crates/secrets-hashicorp/src/secret_manager/raw.rs b/litellm-rust/crates/secrets-hashicorp/src/secret_manager/raw.rs new file mode 100644 index 00000000000..0325561a7a0 --- /dev/null +++ b/litellm-rust/crates/secrets-hashicorp/src/secret_manager/raw.rs @@ -0,0 +1,155 @@ +use super::*; +use rustify::{client::Client as _, endpoint::Endpoint}; +use vaultrs::api::kv2::requests::{ + DeleteLatestSecretVersionRequest, ReadSecretRequest, SetSecretRequest, +}; + +#[derive(veil::Redact)] +pub enum RawOperationError { + Local(Error), + Authentication { + source: Error, + url: String, + certificate: bool, + }, + Http { + method: String, + url: String, + status: u16, + #[redact] + body: Vec, + }, + Transport(#[redact] RustifyClientError), + Timeout { + method: String, + elapsed: Duration, + }, +} + +impl HashicorpVault { + pub async fn write_raw( + &self, + name: &str, + value: &SecretValue, + context: &SecretWriteContext, + ) -> Result, RawOperationError> { + let location = self + .secret_location_with_context(name, &context.operation) + .map_err(RawOperationError::Local)?; + let data = super::write::write_data(value, context).map_err(RawOperationError::Local)?; + let response = self + .raw_request( + &location, + SetSecretRequest { + mount: location.mount.clone(), + path: location.path.clone(), + data, + options: None, + }, + &context.operation, + ) + .await?; + self.cache + .invalidate_where(move |key| key.location == location); + Ok(response) + } + + pub async fn delete_raw( + &self, + name: &str, + context: &HashicorpOperationContext, + ) -> Result<(), RawOperationError> { + let location = self + .secret_location_with_context(name, context) + .map_err(RawOperationError::Local)?; + self.raw_request( + &location, + DeleteLatestSecretVersionRequest { + mount: location.mount.clone(), + path: location.path.clone(), + }, + context, + ) + .await?; + self.cache + .invalidate_where(move |key| key.location == location); + Ok(()) + } + + pub async fn read_raw( + &self, + name: &str, + context: &HashicorpOperationContext, + ) -> Result, RawOperationError> { + let location = self + .secret_location_with_context(name, context) + .map_err(RawOperationError::Local)?; + self.raw_request( + &location, + ReadSecretRequest { + mount: location.mount.clone(), + path: location.path.clone(), + version: None, + }, + context, + ) + .await + } + + async fn raw_request( + &self, + location: &SecretLocation, + endpoint: impl Endpoint, + context: &HashicorpOperationContext, + ) -> Result, RawOperationError> { + let client = self.client_for_location(location).await.map_err(|source| { + let certificate = self.config.approle.is_none(); + let mount = self + .config + .approle + .as_ref() + .map_or("cert", |auth| auth.mount_path.as_str()); + RawOperationError::Authentication { + source, + certificate, + url: format!("{}/v1/auth/{mount}/login", self.config.address), + } + })?; + let request = endpoint + .with_middleware(&client.middle) + .request(client.http.base()) + .map_err(RawOperationError::Transport)?; + let method = request.method().to_string(); + let url = format!( + "{}/v1/{}{}/data/{}", + self.config.address, + location + .namespace + .as_ref() + .map(|ns| format!("{ns}/")) + .unwrap_or_default(), + location.mount, + location.path + ); + let started = Instant::now(); + let response = match context.timeout { + Some(timeout) => tokio::time::timeout(timeout, client.http.send(request)) + .await + .map_err(|_| RawOperationError::Timeout { + method: method.clone(), + elapsed: started.elapsed(), + })?, + None => client.http.send(request).await, + } + .map_err(RawOperationError::Transport)?; + if !response.status().is_success() { + return Err(RawOperationError::Http { + method, + url, + status: response.status().as_u16(), + body: response.into_body(), + }); + } + Ok(response.into_body()) + } +} diff --git a/litellm-rust/crates/secrets-hashicorp/src/secret_manager/read.rs b/litellm-rust/crates/secrets-hashicorp/src/secret_manager/read.rs new file mode 100644 index 00000000000..f88fa46b868 --- /dev/null +++ b/litellm-rust/crates/secrets-hashicorp/src/secret_manager/read.rs @@ -0,0 +1,63 @@ +use super::*; + +impl HashicorpVault { + pub async fn async_read_secret(&self, secret_name: &str) -> Result, Error> { + self.async_read_secret_with_context(secret_name, &HashicorpOperationContext::default()) + .await + } + + pub async fn async_read_secret_with_context( + &self, + secret_name: &str, + context: &HashicorpOperationContext, + ) -> Result, Error> { + let location: SecretLocation = self.secret_location_with_context(secret_name, context)?; + let data_key: String = data_key(context); + let cache_key = CacheKey { + location: location.clone(), + data_key: data_key.clone(), + }; + with_timeout( + context, + self.cache + .read(cache_key, self.read_uncached(&location, &data_key)), + ) + .await + } + + pub(super) async fn read_uncached( + &self, + location: &SecretLocation, + data_key: &str, + ) -> Result, Error> { + let client = self.client_for_location(location).await?; + let data: Option> = + match kv2::read(client.as_ref(), &location.mount, &location.path).await { + Ok(data) => Some(data), + Err(error) if api_status(&error) == Some(404) => None, + Err(error) => return Err(map_api_error(error, ErrorContext::Read)), + }; + let Some(data) = data else { + return Ok(None); + }; + let Some(value) = data.get(data_key) else { + return Ok(None); + }; + let value: &str = value.as_str().ok_or(Error::NonStringValue)?; + let value: SecretValue = SecretValue::new(value); + Ok(Some(value)) + } +} + +impl BaseSecretManager for HashicorpVault { + type Error = Error; + type Context = HashicorpOperationContext; + + async fn async_read_secret( + &self, + name: &str, + context: &Self::Context, + ) -> Result, Error> { + HashicorpVault::async_read_secret_with_context(self, name, context).await + } +} diff --git a/litellm-rust/crates/secrets-hashicorp/src/secret_manager/write.rs b/litellm-rust/crates/secrets-hashicorp/src/secret_manager/write.rs new file mode 100644 index 00000000000..0b1acb3a179 --- /dev/null +++ b/litellm-rust/crates/secrets-hashicorp/src/secret_manager/write.rs @@ -0,0 +1,190 @@ +use super::*; + +impl HashicorpVault { + pub async fn async_write_secret( + &self, + secret_name: &str, + value: SecretValue, + description: Option<&str>, + ) -> Result { + self.async_write_secret_with_context( + secret_name, + &value, + &SecretWriteContext { + description: description.map(str::to_owned), + ..SecretWriteContext::default() + }, + ) + .await + } + + pub async fn async_write_secret_with_context( + &self, + secret_name: &str, + value: &SecretValue, + context: &SecretWriteContext, + ) -> Result { + let location: SecretLocation = + self.secret_location_with_context(secret_name, &context.operation)?; + let data = write_data(value, context)?; + let metadata = with_timeout(&context.operation, async { + let client = self.client_for_location(&location).await?; + match kv2::set(client.as_ref(), &location.mount, &location.path, &data).await { + Ok(metadata) => Ok(metadata), + Err(error) if api_status(&error) == Some(400) => { + let version = + match kv2::read_metadata(client.as_ref(), &location.mount, &location.path) + .await + { + Ok(metadata) => u32::try_from(metadata.current_version) + .map_err(|_| Error::CasVersionOverflow)?, + Err(error) if api_status(&error) == Some(404) => 0, + Err(_) => return Err(map_api_error(error, ErrorContext::Secret)), + }; + kv2::set_with_options( + client.as_ref(), + &location.mount, + &location.path, + &data, + SetSecretRequestOptions { cas: version }, + ) + .await + .map_err(|error| map_api_error(error, ErrorContext::Secret)) + } + Err(error) => Err(map_api_error(error, ErrorContext::Secret)), + } + }) + .await?; + self.cache + .invalidate_where(move |key| key.location == location); + serde_json::to_value(metadata) + .map_err(|source| Error::Client(ClientError::JsonParseError { source })) + } + + pub async fn async_delete_secret(&self, secret_name: &str) -> Result<(), Error> { + self.async_delete_secret_with_context(secret_name, &HashicorpOperationContext::default()) + .await + } + + pub async fn async_delete_secret_with_context( + &self, + secret_name: &str, + context: &HashicorpOperationContext, + ) -> Result<(), Error> { + let location: SecretLocation = self.secret_location_with_context(secret_name, context)?; + with_timeout(context, async { + let client = self.client_for_location(&location).await?; + kv2::delete_latest(client.as_ref(), &location.mount, &location.path) + .await + .map_err(|error| map_api_error(error, ErrorContext::Secret)) + }) + .await?; + self.cache + .invalidate_where(move |key| key.location == location); + Ok(()) + } + + pub async fn async_rotate_secret( + &self, + current_name: &str, + new_name: &str, + value: &SecretValue, + ) -> Result> { + self.async_rotate_secret_with_context( + current_name, + new_name, + value, + &HashicorpOperationContext::default(), + ) + .await + } + + pub async fn async_rotate_secret_with_context( + &self, + current_name: &str, + new_name: &str, + value: &SecretValue, + context: &HashicorpOperationContext, + ) -> Result> { + async_rotate_secret(self, current_name, new_name, value, context).await + } +} + +impl SecretWriter for HashicorpVault { + type WriteResponse = Value; + + async fn async_write_secret( + &self, + name: &str, + value: &SecretValue, + context: &SecretWriteContext, + ) -> Result { + HashicorpVault::async_write_secret_with_context(self, name, value, context).await + } +} + +impl SecretDeleter for HashicorpVault { + type DeleteResponse = (); + + async fn async_delete_secret(&self, name: &str, context: &Self::Context) -> Result<(), Error> { + HashicorpVault::async_delete_secret_with_context(self, name, context).await + } +} + +impl SecretRotator for HashicorpVault { + type RotationResponse = Value; + + async fn async_read_secret_fresh( + &self, + name: &str, + context: &Self::Context, + ) -> Result, Error> { + let location = self.secret_location_with_context(name, context)?; + let data_key = data_key(context); + let key = CacheKey { + location: location.clone(), + data_key: data_key.clone(), + }; + with_timeout( + context, + self.cache + .refresh(key, self.read_uncached(&location, &data_key)), + ) + .await + } + + async fn async_write_replacement( + &self, + current_name: &str, + new_name: &str, + value: &SecretValue, + context: &Self::Context, + ) -> Result { + SecretWriter::async_write_secret( + self, + new_name, + value, + &SecretWriteContext::rotated_from(current_name, context.clone()), + ) + .await + } +} + +pub(super) fn write_data( + value: &SecretValue, + context: &SecretWriteContext, +) -> Result { + let data_key = data_key(&context.operation); + if context.description.is_some() && data_key == "description" { + return Err(Error::DataKeyConflictsWithDescription); + } + let data = std::iter::once((data_key, Value::String(value.expose().to_owned()))) + .chain( + context + .description + .as_ref() + .map(|description| ("description".to_owned(), Value::String(description.clone()))), + ) + .collect(); + Ok(Value::Object(data)) +} diff --git a/litellm-rust/crates/secrets-hashicorp/tests/secret_manager.rs b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager.rs index 6b322de4d8d..668871d3fc7 100644 --- a/litellm-rust/crates/secrets-hashicorp/tests/secret_manager.rs +++ b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager.rs @@ -3,8 +3,8 @@ use std::{collections::HashMap, sync::Arc, time::Duration}; use litellm_core_utils::settings::Lookup; use litellm_secrets_hashicorp::{Error, HashicorpVault, HashicorpVaultConfig}; use litellm_secrets_types::{ - AwsOperationContext, BaseSecretManager, CyberarkOperationContext, HashicorpOperationContext, - SecretOperationContext, SecretValue, SecretWriteContext, + BaseSecretManager, HashicorpOperationContext, RotationError, SecretDeleter, SecretValue, + SecretWriteContext, SecretWriter, }; use rstest::{fixture, rstest}; use serde::Deserialize; @@ -14,923 +14,13 @@ use wiremock::{ matchers::{body_json, header, method, path}, }; -fn config(server: &MockServer, values: &[(&str, &str)]) -> HashicorpVaultConfig { - let mut environment_values: HashMap = values - .iter() - .map(|(name, value)| ((*name).to_owned(), (*value).to_owned())) - .collect(); - environment_values.insert("HCP_VAULT_ADDR".to_owned(), server.uri()); - let environment: Arc = - Arc::new(move |name: &str| environment_values.get(name).cloned()); - HashicorpVaultConfig::from_environment(environment.as_ref()).unwrap() -} +#[path = "secret_manager/support.rs"] +mod support; +use support::*; -fn manager(server: &MockServer, values: &[(&str, &str)]) -> HashicorpVault { - HashicorpVault::from_config(config(server, values), true).unwrap() -} - -fn auth_response(token: &str, lease_duration: u64) -> serde_json::Value { - json!({ - "auth": { - "client_token": token, - "accessor": "", - "policies": [], - "token_policies": [], - "metadata": null, - "lease_duration": lease_duration, - "renewable": false, - "entity_id": "", - "token_type": "service", - "orphan": false - }, - "lease_id": "", - "lease_duration": lease_duration, - "renewable": false, - "request_id": "", - "warnings": null, - "wrap_info": null - }) -} - -fn read_response(data: serde_json::Value) -> serde_json::Value { - json!({ - "data": { - "data": data, - "metadata": { - "created_time": "", - "deletion_time": "", - "custom_metadata": null, - "destroyed": false, - "version": 1 - } - }, - "lease_id": "", - "lease_duration": 0, - "renewable": false, - "request_id": "", - "warnings": null, - "wrap_info": null - }) -} - -#[fixture] -fn token_values() -> Vec<(&'static str, &'static str)> { - vec![("HCP_VAULT_TOKEN", "token")] -} - -#[rstest] -#[tokio::test] -async fn token_reads_use_vault_headers_and_cache_values(token_values: Vec<(&str, &str)>) { - let server: MockServer = MockServer::start().await; - Mock::given(method("GET")) - .and(path("/v1/secret/data/name")) - .and(header("X-Vault-Token", "token")) - .respond_with( - ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), - ) - .expect(1) - .mount(&server) - .await; - let manager: HashicorpVault = manager(&server, &token_values); - - assert_eq!( - manager - .async_read_secret("name") - .await - .unwrap() - .unwrap() - .expose(), - "value" - ); - let requests = server.received_requests().await.unwrap(); - assert!( - requests - .iter() - .all(|request| !request.headers.contains_key("X-Vault-Namespace")) - ); - assert_eq!( - manager - .async_read_secret("name") - .await - .unwrap() - .unwrap() - .expose(), - "value" - ); -} - -#[rstest] -#[tokio::test] -async fn namespace_mount_and_prefix_are_sanitized_in_the_url() { - let server: MockServer = MockServer::start().await; - Mock::given(method("GET")) - .and(path("/v1/kv-prod/data/virtual-keys/name")) - .and(header("X-Vault-Namespace", "team-a")) - .respond_with( - ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), - ) - .expect(1) - .mount(&server) - .await; - let manager: HashicorpVault = manager( - &server, - &[ - ("HCP_VAULT_TOKEN", "token"), - ("HCP_VAULT_SECRET_NAMESPACE", " /team-a/ "), - ("HCP_VAULT_MOUNT_NAME", " /kv-prod/ "), - ("HCP_VAULT_PATH_PREFIX", " /virtual-keys/ "), - ], - ); - - let location = manager.secret_location("name").unwrap(); - assert_eq!(location.namespace.as_deref(), Some("team-a")); - assert_eq!(location.mount, "kv-prod"); - assert_eq!(location.path, "virtual-keys/name"); - assert!(manager.async_read_secret("name").await.unwrap().is_some()); -} - -#[rstest] -fn trailing_address_slashes_are_removed() { - let environment: Arc = Arc::new(|name: &str| match name { - "HCP_VAULT_ADDR" => Some("http://vault.test:8200///".to_owned()), - "HCP_VAULT_TOKEN" => Some("token".to_owned()), - _ => None, - }); - let config: HashicorpVaultConfig = - HashicorpVaultConfig::from_environment(environment.as_ref()).unwrap(); - let manager: HashicorpVault = HashicorpVault::from_config(config, true).unwrap(); - - assert_eq!( - manager.secret_location("name").unwrap(), - litellm_secrets_hashicorp::SecretLocation { - namespace: None, - mount: "secret".to_owned(), - path: "name".to_owned(), - } - ); -} - -#[rstest] -#[case::negative("-1")] -#[case::not_a_number("not-a-number")] -fn invalid_refresh_intervals_are_rejected(#[case] value: &str) { - let environment: Arc = Arc::new(move |name: &str| match name { - "HCP_VAULT_REFRESH_INTERVAL" => Some(value.to_owned()), - _ => None, - }); - - assert!(matches!( - HashicorpVaultConfig::from_environment(environment.as_ref()), - Err(Error::RefreshInterval) - )); -} - -#[rstest] -#[tokio::test] -async fn approle_login_uses_namespace_and_reuses_the_token() { - let server: MockServer = MockServer::start().await; - Mock::given(method("POST")) - .and(path("/v1/auth/custom-approle/login")) - .and(header("X-Vault-Namespace", "login-root")) - .and(body_json(json!({"role_id": "role", "secret_id": "secret"}))) - .respond_with(ResponseTemplate::new(200).set_body_json(auth_response("login-token", 3600))) - .expect(1) - .mount(&server) - .await; - Mock::given(method("GET")) - .and(path("/v1/secret/data/name")) - .and(header("X-Vault-Token", "login-token")) - .and(header("X-Vault-Namespace", "secret-root")) - .respond_with( - ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), - ) - .expect(1) - .mount(&server) - .await; - Mock::given(method("GET")) - .and(path("/v1/secret/data/name-2")) - .respond_with(ResponseTemplate::new(404).set_body_json(json!({"errors": ["missing"]}))) - .expect(1) - .mount(&server) - .await; - let manager: HashicorpVault = manager( - &server, - &[ - ("HCP_VAULT_APPROLE_ROLE_ID", "role"), - ("HCP_VAULT_APPROLE_SECRET_ID", "secret"), - ("HCP_VAULT_APPROLE_MOUNT_PATH", "custom-approle"), - ("HCP_VAULT_NAMESPACE", "secret-root"), - ("HCP_VAULT_LOGIN_NAMESPACE", "login-root"), - ], - ); - - assert!(manager.async_read_secret("name").await.unwrap().is_some()); - assert!(manager.async_read_secret("name-2").await.unwrap().is_none()); -} - -#[rstest] -#[tokio::test] -async fn approle_tokens_expire_after_the_vault_lease() { - let server: MockServer = MockServer::start().await; - Mock::given(method("POST")) - .and(path("/v1/auth/approle/login")) - .respond_with(ResponseTemplate::new(200).set_body_json(auth_response("login-token", 1))) - .expect(2) - .mount(&server) - .await; - Mock::given(method("GET")) - .respond_with( - ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), - ) - .expect(2) - .mount(&server) - .await; - let manager: HashicorpVault = manager( - &server, - &[ - ("HCP_VAULT_APPROLE_ROLE_ID", "role"), - ("HCP_VAULT_APPROLE_SECRET_ID", "secret"), - ("HCP_VAULT_REFRESH_INTERVAL", "0"), - ], - ); - - assert!(manager.async_read_secret("first").await.unwrap().is_some()); - tokio::time::sleep(Duration::from_secs(1) + Duration::from_millis(50)).await; - assert!(manager.async_read_secret("second").await.unwrap().is_some()); -} - -#[rstest] -#[tokio::test] -async fn tls_login_posts_the_role_and_uses_the_client_identity() { - let server: MockServer = MockServer::start().await; - let directory: tempfile::TempDir = tempfile::tempdir().unwrap(); - let cert_path = directory.path().join("client.crt"); - let key_path = directory.path().join("client.key"); - std::fs::write(&cert_path, TEST_CERTIFICATE).unwrap(); - std::fs::write(&key_path, TEST_PRIVATE_KEY).unwrap(); - Mock::given(method("POST")) - .and(path("/v1/auth/cert/login")) - .and(header("X-Vault-Namespace", "login-ns")) - .respond_with(ResponseTemplate::new(200).set_body_json(auth_response("cert-token", 0))) - .expect(2) - .mount(&server) - .await; - Mock::given(method("GET")) - .and(path("/v1/secret/data/name")) - .and(header("X-Vault-Token", "cert-token")) - .and(header("X-Vault-Namespace", "secret-ns")) - .respond_with( - ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), - ) - .expect(2) - .mount(&server) - .await; - let role_values: HashMap = HashMap::from([ - ("HCP_VAULT_ADDR".to_owned(), server.uri()), - ( - "HCP_VAULT_CLIENT_CERT".to_owned(), - cert_path.to_str().unwrap().to_owned(), - ), - ( - "HCP_VAULT_CLIENT_KEY".to_owned(), - key_path.to_str().unwrap().to_owned(), - ), - ("HCP_VAULT_CERT_ROLE".to_owned(), "vault-role".to_owned()), - ( - "HCP_VAULT_LOGIN_NAMESPACE".to_owned(), - "login-ns".to_owned(), - ), - ( - "HCP_VAULT_SECRET_NAMESPACE".to_owned(), - "secret-ns".to_owned(), - ), - ]); - let role_environment: Arc = - Arc::new(move |name: &str| role_values.get(name).cloned()); - let role_manager: HashicorpVault = HashicorpVault::new(role_environment, true).unwrap(); - assert!( - role_manager - .async_read_secret("name") - .await - .unwrap() - .is_some() - ); - - let no_role_values: HashMap = HashMap::from([ - ("HCP_VAULT_ADDR".to_owned(), server.uri()), - ( - "HCP_VAULT_CLIENT_CERT".to_owned(), - cert_path.to_str().unwrap().to_owned(), - ), - ( - "HCP_VAULT_CLIENT_KEY".to_owned(), - key_path.to_str().unwrap().to_owned(), - ), - ( - "HCP_VAULT_LOGIN_NAMESPACE".to_owned(), - "login-ns".to_owned(), - ), - ( - "HCP_VAULT_SECRET_NAMESPACE".to_owned(), - "secret-ns".to_owned(), - ), - ]); - let no_role_environment: Arc = - Arc::new(move |name: &str| no_role_values.get(name).cloned()); - let no_role_manager: HashicorpVault = HashicorpVault::new(no_role_environment, true).unwrap(); - assert!( - no_role_manager - .async_read_secret("name") - .await - .unwrap() - .is_some() - ); - let login_bodies: Vec = server - .received_requests() - .await - .unwrap() - .iter() - .filter(|request| request.method.as_str() == "POST") - .map(|request| serde_json::from_slice(&request.body).unwrap()) - .collect(); - assert!(login_bodies.contains(&json!({"name": "vault-role"}))); - assert!(login_bodies.contains(&json!({}))); -} - -#[derive(Clone, Copy)] -enum ExpectedRead { - Missing, - Malformed, - NonString, -} - -#[rstest] -#[case::missing(404, json!({"errors": ["missing"]}), ExpectedRead::Missing)] -#[case::malformed(200, json!({"data": "invalid"}), ExpectedRead::Malformed)] -#[case::missing_key(200, json!({}), ExpectedRead::Missing)] -#[case::non_string(200, json!({"key": 1}), ExpectedRead::NonString)] -#[tokio::test] -async fn read_responses_distinguish_absence_and_malformed_payloads( - token_values: Vec<(&str, &str)>, - #[case] status: u16, - #[case] body: serde_json::Value, - #[case] expected: ExpectedRead, -) { - let server: MockServer = MockServer::start().await; - Mock::given(method("GET")) - .respond_with(ResponseTemplate::new(status).set_body_json( - if status == 200 && !matches!(expected, ExpectedRead::Malformed) { - read_response(body) - } else { - body - }, - )) - .expect(1) - .mount(&server) - .await; - let result: Result, Error> = manager(&server, &token_values) - .async_read_secret("name") - .await; - match expected { - ExpectedRead::Missing => assert!(result.unwrap().is_none()), - ExpectedRead::Malformed => assert!(matches!(result, Err(Error::MalformedPayload))), - ExpectedRead::NonString => assert!(matches!(result, Err(Error::NonStringValue))), - } -} - -#[rstest] -#[tokio::test] -async fn write_and_delete_invalidate_the_read_cache(token_values: Vec<(&str, &str)>) { - let server: MockServer = MockServer::start().await; - Mock::given(method("GET")) - .and(path("/v1/secret/data/name")) - .respond_with( - ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), - ) - .expect(2) - .mount(&server) - .await; - Mock::given(method("POST")) - .and(path("/v1/secret/data/name")) - .and(body_json( - json!({"data": {"key": "updated", "description": "description"}}), - )) - .respond_with(ResponseTemplate::new(200).set_body_json(json!({ - "data": { - "created_time": "", - "deletion_time": "", - "custom_metadata": null, - "destroyed": false, - "version": 2 - }, - "lease_id": "", - "lease_duration": 0, - "renewable": false, - "request_id": "", - "warnings": null, - "wrap_info": null - }))) - .expect(1) - .mount(&server) - .await; - Mock::given(method("DELETE")) - .and(path("/v1/secret/data/name")) - .respond_with(ResponseTemplate::new(204)) - .expect(1) - .mount(&server) - .await; - let manager: HashicorpVault = manager(&server, &token_values); - - assert!(manager.async_read_secret("name").await.unwrap().is_some()); - assert!( - manager - .async_write_secret("name", SecretValue::new("updated"), Some("description")) - .await - .is_ok() - ); - assert!(manager.async_read_secret("name").await.unwrap().is_some()); - manager.async_delete_secret("name").await.unwrap(); -} - -#[rstest] -#[tokio::test] -async fn base_manager_context_overrides_vault_location_and_data_key( - token_values: Vec<(&str, &str)>, -) { - let server: MockServer = MockServer::start().await; - Mock::given(method("GET")) - .and(path("/v1/alternate/data/managed/name")) - .respond_with( - ResponseTemplate::new(200).set_body_json(read_response(json!({"api_token": "value"}))), - ) - .expect(2) - .mount(&server) - .await; - Mock::given(method("POST")) - .and(path("/v1/alternate/data/managed/name")) - .and(body_json(json!({ - "data": {"api_token": "updated", "description": "Managed key"} - }))) - .respond_with(ResponseTemplate::new(200).set_body_json(json!({ - "data": { - "created_time": "", - "deletion_time": "", - "custom_metadata": null, - "destroyed": false, - "version": 2 - }, - "lease_id": "", - "lease_duration": 0, - "renewable": false, - "request_id": "", - "warnings": null, - "wrap_info": null - }))) - .expect(1) - .mount(&server) - .await; - Mock::given(method("DELETE")) - .and(path("/v1/alternate/data/managed/name")) - .respond_with(ResponseTemplate::new(204)) - .expect(1) - .mount(&server) - .await; - let manager: HashicorpVault = manager(&server, &token_values); - let operation = SecretOperationContext::Hashicorp(HashicorpOperationContext { - mount: Some(" /alternate/ ".to_owned()), - path_prefix: Some(" /managed/ ".to_owned()), - data_key: Some("api_token".to_owned()), - ..HashicorpOperationContext::default() - }); - let write_context = SecretWriteContext { - description: Some("Managed key".to_owned()), - operation: operation.clone(), - ..SecretWriteContext::default() - }; - - assert_eq!( - BaseSecretManager::async_read_secret(&manager, "name", &operation) - .await - .unwrap() - .unwrap() - .expose(), - "value" - ); - BaseSecretManager::async_write_secret( - &manager, - "name", - &SecretValue::new("updated"), - &write_context, - ) - .await - .unwrap(); - assert!( - BaseSecretManager::async_read_secret(&manager, "name", &operation) - .await - .unwrap() - .is_some() - ); - BaseSecretManager::async_delete_secret(&manager, "name", None, &operation) - .await - .unwrap(); -} - -#[rstest] -#[tokio::test] -async fn reads_cache_each_data_key_for_the_same_vault_path(token_values: Vec<(&str, &str)>) { - let server: MockServer = MockServer::start().await; - Mock::given(method("GET")) - .and(path("/v1/secret/data/name")) - .respond_with( - ResponseTemplate::new(200).set_body_json(read_response(json!({ - "key": "primary", - "alternate": "secondary" - }))), - ) - .expect(2) - .mount(&server) - .await; - let manager: HashicorpVault = manager(&server, &token_values); - let alternate = SecretOperationContext::Hashicorp(HashicorpOperationContext { - data_key: Some("alternate".to_owned()), - ..HashicorpOperationContext::default() - }); - - assert_eq!( - manager - .async_read_secret("name") - .await - .unwrap() - .unwrap() - .expose(), - "primary" - ); - assert_eq!( - BaseSecretManager::async_read_secret(&manager, "name", &alternate) - .await - .unwrap() - .unwrap() - .expose(), - "secondary" - ); - assert_eq!( - manager - .async_read_secret("name") - .await - .unwrap() - .unwrap() - .expose(), - "primary" - ); -} - -#[rstest] -#[tokio::test] -async fn base_manager_context_timeout_limits_vault_io(token_values: Vec<(&str, &str)>) { - let server: MockServer = MockServer::start().await; - Mock::given(method("GET")) - .and(path("/v1/secret/data/name")) - .respond_with( - ResponseTemplate::new(200) - .set_delay(Duration::from_millis(100)) - .set_body_json(read_response(json!({"key": "value"}))), - ) - .expect(1) - .mount(&server) - .await; - let manager: HashicorpVault = manager(&server, &token_values); - let context = SecretOperationContext::Hashicorp(HashicorpOperationContext { - timeout: Some(Duration::from_millis(10)), - ..HashicorpOperationContext::default() - }); - - assert!(matches!( - BaseSecretManager::async_read_secret(&manager, "name", &context).await, - Err(Error::Timeout) - )); -} - -#[rstest] -#[case::aws(SecretOperationContext::Aws(AwsOperationContext::default()))] -#[case::cyberark(SecretOperationContext::Cyberark(CyberarkOperationContext::default()))] -#[tokio::test] -async fn foreign_contexts_cannot_access_vault_secrets( - token_values: Vec<(&str, &str)>, - #[case] context: SecretOperationContext, - #[values(false, true)] cached: bool, -) { - let server = MockServer::start().await; - Mock::given(method("GET")) - .and(path("/v1/secret/data/name")) - .respond_with( - ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), - ) - .expect(u64::from(cached)) - .mount(&server) - .await; - let manager = manager(&server, &token_values); - if cached { - assert!(manager.async_read_secret("name").await.unwrap().is_some()); - } - assert!(matches!( - BaseSecretManager::async_read_secret(&manager, "name", &context).await, - Err(Error::InvalidOperationContext) - )); - assert!(matches!( - BaseSecretManager::async_write_secret( - &manager, - "name", - &SecretValue::new("replacement"), - &SecretWriteContext { - operation: context.clone(), - ..SecretWriteContext::default() - }, - ) - .await, - Err(Error::InvalidOperationContext) - )); - assert!(matches!( - BaseSecretManager::async_delete_secret(&manager, "name", None, &context).await, - Err(Error::InvalidOperationContext) - )); - assert!(matches!( - manager - .async_rotate_secret_with_context( - "name", - "new", - &SecretValue::new("replacement"), - &context - ) - .await, - Err(Error::InvalidOperationContext) - )); - assert_eq!( - server.received_requests().await.unwrap().len(), - usize::from(cached) - ); -} - -#[rstest] -#[tokio::test] -async fn rotation_applies_timeout_to_each_request(token_values: Vec<(&str, &str)>) { - let server = MockServer::start().await; - let timeout = Duration::from_secs(1); - let delay = timeout / 2; - Mock::given(method("GET")) - .and(path("/v1/alternate/data/managed/current")) - .respond_with( - ResponseTemplate::new(200) - .set_delay(delay) - .set_body_json(read_response(json!({"api_token": "original"}))), - ) - .mount(&server) - .await; - Mock::given(method("POST")) - .and(path("/v1/alternate/data/managed/new")) - .and(body_json(json!({ - "data": {"api_token": "replacement", "description": "Rotated from current"} - }))) - .respond_with( - ResponseTemplate::new(200) - .set_delay(delay) - .set_body_json(json!({ - "data": { - "created_time": "", - "deletion_time": "", - "custom_metadata": null, - "destroyed": false, - "version": 1 - }, - "lease_id": "", - "lease_duration": 0, - "renewable": false, - "request_id": "", - "warnings": null, - "wrap_info": null - })), - ) - .mount(&server) - .await; - Mock::given(method("GET")) - .and(path("/v1/alternate/data/managed/new")) - .respond_with( - ResponseTemplate::new(200) - .set_delay(delay) - .set_body_json(read_response(json!({"api_token": "replacement"}))), - ) - .mount(&server) - .await; - Mock::given(method("DELETE")) - .and(path("/v1/alternate/data/managed/current")) - .respond_with(ResponseTemplate::new(204).set_delay(delay)) - .mount(&server) - .await; - let manager = manager(&server, &token_values); - let context = SecretOperationContext::Hashicorp(HashicorpOperationContext { - timeout: Some(timeout), - mount: Some("alternate".to_owned()), - path_prefix: Some("managed".to_owned()), - data_key: Some("api_token".to_owned()), - }); - - manager - .async_rotate_secret_with_context( - "current", - "new", - &SecretValue::new("replacement"), - &context, - ) - .await - .unwrap(); - let requests = server.received_requests().await.unwrap(); - let operations: Vec<_> = requests - .iter() - .map(|request| (request.method.as_str(), request.url.path())) - .collect(); - assert_eq!( - operations, - [ - ("GET", "/v1/alternate/data/managed/current"), - ("POST", "/v1/alternate/data/managed/new"), - ("GET", "/v1/alternate/data/managed/new"), - ("DELETE", "/v1/alternate/data/managed/current"), - ] - ); -} - -#[rstest] -#[tokio::test] -async fn no_auth_and_invalid_names_fail_without_requests() { - let server: MockServer = MockServer::start().await; - let manager: HashicorpVault = manager(&server, &[]); - - assert!(matches!( - manager.async_read_secret("name").await, - Err(Error::NoAuthConfigured) - )); - assert!(matches!( - manager.async_read_secret("../name").await, - Err(Error::InvalidSecretName(_)) - )); - assert!(server.received_requests().await.unwrap().is_empty()); -} - -#[rstest] -#[tokio::test] -async fn debug_output_redacts_authentication_values() { - let server: MockServer = MockServer::start().await; - let manager: HashicorpVault = - HashicorpVault::from_config(config(&server, &[("HCP_VAULT_TOKEN", "token-value")]), true) - .unwrap(); - let debug: String = format!("{manager:?}"); - assert!(!debug.contains("token-value")); - assert!(!debug.contains("secret-id")); -} - -#[derive(Deserialize)] -struct ParityCase { - env: HashMap, - expected_secret_url: String, - expected_login_url: Option, - expected_login_namespace: Option, - expected_secret_namespace: Option, - secret_name: String, -} - -#[fixture] -fn parity_cases() -> Vec { - serde_json::from_str(include_str!(concat!( - env!("CARGO_MANIFEST_DIR"), - "/../../../tests/test_litellm/secret_managers/hashicorp_vault_parity.json" - ))) - .unwrap() -} - -#[rstest] -fn configuration_matches_python_parity_fixture(parity_cases: Vec) { - for case in parity_cases { - let values: HashMap = case.env.clone(); - let environment: Arc = - Arc::new(move |name: &str| values.get(name).cloned()); - let config: HashicorpVaultConfig = - HashicorpVaultConfig::from_environment(environment.as_ref()).unwrap(); - let manager: HashicorpVault = HashicorpVault::from_config(config.clone(), true).unwrap(); - let location = manager.secret_location(&case.secret_name).unwrap(); - let namespace = location - .namespace - .as_deref() - .map(|namespace| format!("{namespace}/")) - .unwrap_or_default(); - assert_eq!( - format!( - "{}/v1/{}{}/data/{}", - config.address, namespace, location.mount, location.path - ), - case.expected_secret_url - ); - let login_url = config.approle.as_ref().map_or_else( - || { - config - .tls_cert - .as_ref() - .map(|_| format!("{}/v1/auth/cert/login", config.address)) - }, - |approle| { - Some(format!( - "{}/v1/auth/{}/login", - config.address, approle.mount_path - )) - }, - ); - assert_eq!(login_url, case.expected_login_url); - assert_eq!( - manager.config().login_namespace(), - case.expected_login_namespace.as_deref() - ); - assert_eq!( - manager.config().secret_namespace(), - case.expected_secret_namespace.as_deref() - ); - } -} - -#[rstest] -#[tokio::test] -#[ignore] -async fn live_vault_round_trip() { - let environment: Arc = - Arc::new(litellm_core_utils::settings::ProcessEnvironment); - let manager: HashicorpVault = HashicorpVault::new(environment, true).unwrap(); - let name: String = std::env::var("LITELLM_VAULT_LIVE_SECRET_NAME").unwrap(); - let value: SecretValue = SecretValue::new("native-live-value"); - let location = manager.secret_location(&name).unwrap(); - println!( - "native provenance: {} vaultrs {} {:?} {} {}", - module_path!(), - manager.config().address, - location.namespace, - location.mount, - location.path - ); - manager - .async_write_secret(&name, value.clone(), None) - .await - .unwrap(); - assert_eq!( - manager.async_read_secret(&name).await.unwrap().unwrap(), - value - ); - manager.async_delete_secret(&name).await.unwrap(); - assert!(manager.async_read_secret(&name).await.unwrap().is_none()); -} - -const TEST_CERTIFICATE: &str = "-----BEGIN CERTIFICATE----- -MIIDDzCCAfegAwIBAgIUeMzLFLM/mRbPGbNAew5N2UTscocwDQYJKoZIhvcNAQEL -BQAwFzEVMBMGA1UEAwwMbGl0ZWxsbS10ZXN0MB4XDTI2MDkyMTIwMjA1OVoXDTI2 -MDkyMjIwMjA1OVowFzEVMBMGA1UEAwwMbGl0ZWxsbS10ZXN0MIIBIjANBgkqhkiG -9w0BAQEFAAOCAQ8AMIIBCgKCAQEAveYoSUJXybmkHmQsBfhBcv2Ob5Oy8ejZu+B3 -vTnrPumW4ANi1XXKBSazRGB3fEtAgr+3KhKeHaSKEQeBwJkAEBfdmQv0tpXICwHs -1kFNtU0owy54HVW5/ia+LMszsFcPzVIoMnbUOuiKr9RaV7P+IEFzILPBVuV4DoYH -yocjD3+9QNqokWgNL8LK37JijmNEFVaKFz0X6SyL2VRDlfPWTEBK52Gp/pvDgA6G -eTSfyI+kCm9h5ECTYUAtmatk9WPVS8sWOqV1EXVanFyYBU+mDxoywAS1/6CHeIPh -bNmCOZjPoO9qWBJ7ZyGhOconBigXY8qnlXymev+44IPHrx4urwIDAQABo1MwUTAd -BgNVHQ4EFgQUvaZrZ6HKtbr3ekeZmgy4b5Pq95QwHwYDVR0jBBgwFoAUvaZrZ6HK -tbr3ekeZmgy4b5Pq95QwDwYDVR0TAQH/BAUwAwEB/zANBgkqhkiG9w0BAQsFAAOC -AQEAEejrD8d1qDxW55XxQ4IC31rufoEvDV955jyvh2kALPaN/i5oWsBGI+UAQZna -aaoQXwzlmHrtDUBWl0LztVTUamIleUep2+PLLauqqt43vxppxMX8Jn2mnPO20YE/ -hIzGx0jN/LBG8PDyLSvHdlgjP9ofA4Vg4rTQugdXRgOvlCE/epnH/MADcg9KYJtJ -C1RObCIkL3LcdUbjStJRCY/U/FeWcgyncEPz95OFDkbrlNDajb6o6CkYfouqvhTc -8XlgjjAVKIbAbRgbVu3elsquuFM97x2DzWDjkrMNmDt1FJ9ubK36gL6B3o0UMaoQ -00R7x/eqvH+EkWa/2ekW9lpleQ== ------END CERTIFICATE----- -"; - -const TEST_PRIVATE_KEY: &str = "-----BEGIN PRIVATE KEY----- -MIIEvQIBADANBgkqhkiG9w0BAQEFAASCBKcwggSjAgEAAoIBAQC95ihJQlfJuaQe -ZCwF+EFy/Y5vk7Lx6Nm74He9Oes+6ZbgA2LVdcoFJrNEYHd8S0CCv7cqEp4dpIoR -B4HAmQAQF92ZC/S2lcgLAezWQU21TSjDLngdVbn+Jr4syzOwVw/NUigydtQ66Iqv -1FpXs/4gQXMgs8FW5XgOhgfKhyMPf71A2qiRaA0vwsrfsmKOY0QVVooXPRfpLIvZ -VEOV89ZMQErnYan+m8OADoZ5NJ/Ij6QKb2HkQJNhQC2Zq2T1Y9VLyxY6pXURdVqc -XJgFT6YPGjLABLX/oId4g+Fs2YI5mM+g72pYEntnIaE5yicGKBdjyqeVfKZ6/7jg -g8evHi6vAgMBAAECggEAGdJjlP6b8Fa5bdaCM/ebcrbuuNZVJVbb0JPHxGfNSLs7 -pE9hj5QaOdQW2Uviw3h6F61ZCzQH4xD+Iy2po5ZKb2XHYKnDB1bboj+LRGER337T -9aJqe9at2VTMVEv3Rdm40NsEk0QcPLxlK16NQFK90gYEUSSQPDAswJDSG2R/zHn+ -vADI907mW/goEJHeLn8PWGlNlSiR6x+5JJtq+GXCzUzVvJYQSCLGxCSl2x2H+0g7 -NhFI0zPpdzNmO/h+yhzaFb6Rp5U8+ZsnZ3qYjQ/03gw1myTDKJt1YaO9JvArnNYX -hcJQQ8Rt0bHhcrZA16bBOpqZlo5pKCicwI/netgN8QKBgQDcFz7AzdJ26sMSV32V -rwrMgIoggt8qDjO1ARwqW35A1TIge0FoW4M4KpsXQGGfT341uU1esXEcyZ/1L/5X -3ql2gX4DbOYLZLWYzZGR2hq33oi8HkhN98QrEwL9emSH8NqYX3Xxja3PrmCrSYJe -Zbnd9TIm2XkxyMoyXJu6M/QvnwKBgQDc4dzqTbxoGEGa5MuJoGmMwPnqgdG9UM5J -eExVnh7osxc2sOdsiPeRjjQTxs9v2kJwctC359OJoo9yGaaJeSghU4LEWJo1sqnA -fzSCLammYvtVAtniyNv5Mxk/6Uimi4NNDKaAKB+m4K2uSn3U9AmY7KPYMGaSbS9W -XSnobjxm8QKBgC8bPpAvvWs8ZhIn7bY659nLbUT2HeO3dHO6UBf0yzn/J6JyHxbB -93zvCZDZc8uQTRgcmCW7XtVlhjoJUqvl+Wlm39zF0xr/LCsPXKfWAb/2/lcdOCaP -8Emz4QD10EyUTYUtcWYJB/mafhBLRH8F0Nlj4J8WDu2L51MOJTqeYhZLAoGAWffN -icocAbJPlo22sdoa4+/+W5yBF8GAJMDRJtZ+9H1t6SLpQHYRkMIBSETkXUTjZvX9 -Ocs9iIQkNW9pO/mTdO+VBfCo71JUfknR02xR+6m5gYjlws/ZeYlssXGN2/hbhNiw -QOcW7Vv6olFJK6Iy/oz0t6wPO3kpnN3Zogi0paECgYEAwo44M1DdYCtV0snhmYM9 -5u0mPfYt5P2SVLXyUbr+vFTfrTL/WKnXIJgbsnj3Gvf+GIZv9tKcXhSNmEHQCYX4 -X3w9iTPddCHuvZ1fpufi2TyArJh0OkoNtLXJHTKrHjf2N+61AQzFiv5WieJrdE+H -qr32PTUuVGPyO9LyTY4/RL0= ------END PRIVATE KEY----- -"; +#[path = "secret_manager/configuration.rs"] +mod configuration; +#[path = "secret_manager/reads.rs"] +mod reads; +#[path = "secret_manager/writes.rs"] +mod writes; diff --git a/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/configuration.rs b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/configuration.rs new file mode 100644 index 00000000000..be3059ce986 --- /dev/null +++ b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/configuration.rs @@ -0,0 +1,335 @@ +use super::*; + +#[rstest] +fn trailing_address_slashes_are_removed() { + let environment: Arc = Arc::new(|name: &str| match name { + "HCP_VAULT_ADDR" => Some("http://vault.test:8200///".to_owned()), + "HCP_VAULT_TOKEN" => Some("token".to_owned()), + _ => None, + }); + let config: HashicorpVaultConfig = + HashicorpVaultConfig::from_environment(environment.as_ref()).unwrap(); + let manager: HashicorpVault = HashicorpVault::from_config(config, true).unwrap(); + + assert_eq!( + manager.secret_location("name").unwrap(), + litellm_secrets_hashicorp::SecretLocation { + namespace: None, + mount: "secret".to_owned(), + path: "name".to_owned(), + } + ); +} + +#[rstest] +#[case::negative("-1")] +#[case::not_a_number("not-a-number")] +fn invalid_refresh_intervals_are_rejected(#[case] value: &str) { + let environment: Arc = Arc::new(move |name: &str| match name { + "HCP_VAULT_REFRESH_INTERVAL" => Some(value.to_owned()), + _ => None, + }); + + assert!(matches!( + HashicorpVaultConfig::from_environment(environment.as_ref()), + Err(Error::RefreshInterval) + )); +} + +#[rstest] +#[tokio::test] +async fn approle_login_uses_namespace_and_reuses_the_token() { + let server: MockServer = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/auth/custom-approle/login")) + .and(header("X-Vault-Namespace", "login-root")) + .and(body_json(json!({"role_id": "role", "secret_id": "secret"}))) + .respond_with(ResponseTemplate::new(200).set_body_json(auth_response("login-token", 3600))) + .expect(1) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/name")) + .and(header("X-Vault-Token", "login-token")) + .and(header("X-Vault-Namespace", "secret-root")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), + ) + .expect(1) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/name-2")) + .respond_with(ResponseTemplate::new(404).set_body_json(json!({"errors": ["missing"]}))) + .expect(1) + .mount(&server) + .await; + let manager: HashicorpVault = manager( + &server, + &[ + ("HCP_VAULT_APPROLE_ROLE_ID", "role"), + ("HCP_VAULT_APPROLE_SECRET_ID", "secret"), + ("HCP_VAULT_APPROLE_MOUNT_PATH", "custom-approle"), + ("HCP_VAULT_NAMESPACE", "secret-root"), + ("HCP_VAULT_LOGIN_NAMESPACE", "login-root"), + ], + ); + + assert!(manager.async_read_secret("name").await.unwrap().is_some()); + assert!(manager.async_read_secret("name-2").await.unwrap().is_none()); +} + +#[rstest] +#[tokio::test] +async fn tls_login_posts_the_role_and_uses_the_client_identity() { + let server: MockServer = MockServer::start().await; + let directory: tempfile::TempDir = tempfile::tempdir().unwrap(); + let cert_path = directory.path().join("client.crt"); + let key_path = directory.path().join("client.key"); + std::fs::write(&cert_path, TEST_CERTIFICATE).unwrap(); + std::fs::write(&key_path, TEST_PRIVATE_KEY).unwrap(); + Mock::given(method("POST")) + .and(path("/v1/auth/cert/login")) + .and(header("X-Vault-Namespace", "login-ns")) + .respond_with(ResponseTemplate::new(200).set_body_json(auth_response("cert-token", 0))) + .expect(2) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/name")) + .and(header("X-Vault-Token", "cert-token")) + .and(header("X-Vault-Namespace", "secret-ns")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), + ) + .expect(2) + .mount(&server) + .await; + let role_values: HashMap = HashMap::from([ + ("HCP_VAULT_ADDR".to_owned(), server.uri()), + ( + "HCP_VAULT_CLIENT_CERT".to_owned(), + cert_path.to_str().unwrap().to_owned(), + ), + ( + "HCP_VAULT_CLIENT_KEY".to_owned(), + key_path.to_str().unwrap().to_owned(), + ), + ("HCP_VAULT_CERT_ROLE".to_owned(), "vault-role".to_owned()), + ( + "HCP_VAULT_LOGIN_NAMESPACE".to_owned(), + "login-ns".to_owned(), + ), + ( + "HCP_VAULT_SECRET_NAMESPACE".to_owned(), + "secret-ns".to_owned(), + ), + ]); + let role_environment: Arc = + Arc::new(move |name: &str| role_values.get(name).cloned()); + let role_manager: HashicorpVault = HashicorpVault::new(role_environment, true).unwrap(); + assert!( + role_manager + .async_read_secret("name") + .await + .unwrap() + .is_some() + ); + + let no_role_values: HashMap = HashMap::from([ + ("HCP_VAULT_ADDR".to_owned(), server.uri()), + ( + "HCP_VAULT_CLIENT_CERT".to_owned(), + cert_path.to_str().unwrap().to_owned(), + ), + ( + "HCP_VAULT_CLIENT_KEY".to_owned(), + key_path.to_str().unwrap().to_owned(), + ), + ( + "HCP_VAULT_LOGIN_NAMESPACE".to_owned(), + "login-ns".to_owned(), + ), + ( + "HCP_VAULT_SECRET_NAMESPACE".to_owned(), + "secret-ns".to_owned(), + ), + ]); + let no_role_environment: Arc = + Arc::new(move |name: &str| no_role_values.get(name).cloned()); + let no_role_manager: HashicorpVault = HashicorpVault::new(no_role_environment, true).unwrap(); + assert!( + no_role_manager + .async_read_secret("name") + .await + .unwrap() + .is_some() + ); + let login_bodies: Vec = server + .received_requests() + .await + .unwrap() + .iter() + .filter(|request| request.method.as_str() == "POST") + .map(|request| serde_json::from_slice(&request.body).unwrap()) + .collect(); + assert!(login_bodies.contains(&json!({"name": "vault-role"}))); + assert!(login_bodies.contains(&json!({}))); +} + +#[rstest] +#[tokio::test] +async fn no_auth_and_invalid_names_fail_without_requests() { + let server: MockServer = MockServer::start().await; + let manager: HashicorpVault = manager(&server, &[]); + + assert!(matches!( + manager.async_read_secret("name").await, + Err(Error::NoAuthConfigured) + )); + assert!(matches!( + manager.async_read_secret("../name").await, + Err(Error::InvalidSecretName(_)) + )); + assert!(server.received_requests().await.unwrap().is_empty()); +} + +#[rstest] +#[tokio::test] +async fn debug_output_redacts_authentication_values() { + let server: MockServer = MockServer::start().await; + let manager: HashicorpVault = + HashicorpVault::from_config(config(&server, &[("HCP_VAULT_TOKEN", "token-value")]), true) + .unwrap(); + let debug: String = format!("{manager:?}"); + assert!(!debug.contains("token-value")); + assert!(!debug.contains("secret-id")); +} + +#[rstest] +fn configuration_matches_python_parity_fixture(parity_cases: Vec) { + for case in parity_cases { + let values: HashMap = case.env.clone(); + let environment: Arc = + Arc::new(move |name: &str| values.get(name).cloned()); + let config: HashicorpVaultConfig = + HashicorpVaultConfig::from_environment(environment.as_ref()).unwrap(); + let manager: HashicorpVault = HashicorpVault::from_config(config.clone(), true).unwrap(); + let location = manager.secret_location(&case.secret_name).unwrap(); + let namespace = location + .namespace + .as_deref() + .map(|namespace| format!("{namespace}/")) + .unwrap_or_default(); + assert_eq!( + format!( + "{}/v1/{}{}/data/{}", + config.address, namespace, location.mount, location.path + ), + case.expected_secret_url + ); + let login_url = config.approle.as_ref().map_or_else( + || { + config + .tls_cert + .as_ref() + .map(|_| format!("{}/v1/auth/cert/login", config.address)) + }, + |approle| { + Some(format!( + "{}/v1/auth/{}/login", + config.address, approle.mount_path + )) + }, + ); + assert_eq!(login_url, case.expected_login_url); + assert_eq!( + manager.config().login_namespace(), + case.expected_login_namespace.as_deref() + ); + assert_eq!( + manager.config().secret_namespace(), + case.expected_secret_namespace.as_deref() + ); + } +} + +#[rstest] +#[case::separate( + Some("legacy"), + Some("root"), + Some("teams/team-a"), + Some("root"), + Some("teams/team-a") +)] +#[case::legacy(Some("admin"), None, None, Some("admin"), Some("admin"))] +#[case::login_override(Some("admin"), Some("root"), None, Some("root"), Some("admin"))] +#[case::secret_override( + Some("admin"), + None, + Some("teams/team-a"), + Some("admin"), + Some("teams/team-a") +)] +#[case::no_namespace(None, None, None, None, None)] +#[tokio::test] +async fn login_and_secret_namespaces_follow_python_precedence( + #[case] legacy: Option<&str>, + #[case] login: Option<&str>, + #[case] secret: Option<&str>, + #[case] expected_login: Option<&'static str>, + #[case] expected_secret: Option<&'static str>, +) { + let server = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/auth/approle/login")) + .and(body_json(json!({"role_id":"role", "secret_id":"secret"}))) + .respond_with(move |request: &wiremock::Request| { + assert_eq!( + request + .headers + .get("X-Vault-Namespace") + .map(|value| value.to_str().unwrap()), + expected_login + ); + ResponseTemplate::new(200).set_body_json(auth_response("login-token", 3600)) + }) + .expect(1) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/key")) + .and(header("X-Vault-Token", "login-token")) + .respond_with(move |request: &wiremock::Request| { + assert_eq!( + request + .headers + .get("X-Vault-Namespace") + .map(|value| value.to_str().unwrap()), + expected_secret + ); + ResponseTemplate::new(200).set_body_json(read_response(json!({"key":"value"}))) + }) + .expect(1) + .mount(&server) + .await; + let values: Vec<_> = [ + ("HCP_VAULT_NAMESPACE", legacy), + ("HCP_VAULT_LOGIN_NAMESPACE", login), + ("HCP_VAULT_SECRET_NAMESPACE", secret), + ("HCP_VAULT_APPROLE_ROLE_ID", Some("role")), + ("HCP_VAULT_APPROLE_SECRET_ID", Some("secret")), + ] + .into_iter() + .filter_map(|(key, value)| value.map(|value| (key, value))) + .collect(); + assert_eq!( + manager(&server, &values) + .async_read_secret("key") + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); +} diff --git a/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/reads.rs b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/reads.rs new file mode 100644 index 00000000000..0c662eb5b55 --- /dev/null +++ b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/reads.rs @@ -0,0 +1,385 @@ +use super::*; + +#[rstest] +#[tokio::test] +async fn token_reads_use_vault_headers_and_cache_values(token_values: Vec<(&str, &str)>) { + let server: MockServer = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/name")) + .and(header("X-Vault-Token", "token")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), + ) + .expect(1) + .mount(&server) + .await; + let manager: HashicorpVault = manager(&server, &token_values); + + assert_eq!( + manager + .async_read_secret("name") + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); + let requests = server.received_requests().await.unwrap(); + assert!( + requests + .iter() + .all(|request| !request.headers.contains_key("X-Vault-Namespace")) + ); + assert_eq!( + manager + .async_read_secret("name") + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); +} + +#[rstest] +#[tokio::test] +async fn namespace_mount_and_prefix_are_sanitized_in_the_url() { + let server: MockServer = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/kv-prod/data/virtual-keys/name")) + .and(header("X-Vault-Namespace", "team-a")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), + ) + .expect(1) + .mount(&server) + .await; + let manager: HashicorpVault = manager( + &server, + &[ + ("HCP_VAULT_TOKEN", "token"), + ("HCP_VAULT_SECRET_NAMESPACE", " /team-a/ "), + ("HCP_VAULT_MOUNT_NAME", " /kv-prod/ "), + ("HCP_VAULT_PATH_PREFIX", " /virtual-keys/ "), + ], + ); + + let location = manager.secret_location("name").unwrap(); + assert_eq!(location.namespace.as_deref(), Some("team-a")); + assert_eq!(location.mount, "kv-prod"); + assert_eq!(location.path, "virtual-keys/name"); + assert!(manager.async_read_secret("name").await.unwrap().is_some()); +} + +#[rstest] +#[tokio::test] +async fn approle_tokens_expire_after_the_vault_lease() { + let server: MockServer = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/auth/approle/login")) + .respond_with(ResponseTemplate::new(200).set_body_json(auth_response("login-token", 1))) + .expect(2) + .mount(&server) + .await; + Mock::given(method("GET")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "value"}))), + ) + .expect(2) + .mount(&server) + .await; + let manager: HashicorpVault = manager( + &server, + &[ + ("HCP_VAULT_APPROLE_ROLE_ID", "role"), + ("HCP_VAULT_APPROLE_SECRET_ID", "secret"), + ("HCP_VAULT_REFRESH_INTERVAL", "0"), + ], + ); + + assert!(manager.async_read_secret("first").await.unwrap().is_some()); + tokio::time::sleep(Duration::from_secs(1) + Duration::from_millis(50)).await; + assert!(manager.async_read_secret("second").await.unwrap().is_some()); +} + +#[derive(Clone, Copy)] +enum ExpectedRead { + Missing, + Malformed, + NonString, +} + +#[rstest] +#[case::missing(404, json!({"errors": ["missing"]}), ExpectedRead::Missing)] +#[case::malformed(200, json!({"data": "invalid"}), ExpectedRead::Malformed)] +#[case::missing_key(200, json!({}), ExpectedRead::Missing)] +#[case::non_string(200, json!({"key": 1}), ExpectedRead::NonString)] +#[tokio::test] +async fn read_responses_distinguish_absence_and_malformed_payloads( + token_values: Vec<(&str, &str)>, + #[case] status: u16, + #[case] body: serde_json::Value, + #[case] expected: ExpectedRead, +) { + let server: MockServer = MockServer::start().await; + Mock::given(method("GET")) + .respond_with(ResponseTemplate::new(status).set_body_json( + if status == 200 && !matches!(expected, ExpectedRead::Malformed) { + read_response(body) + } else { + body + }, + )) + .expect(1) + .mount(&server) + .await; + let result: Result, Error> = manager(&server, &token_values) + .async_read_secret("name") + .await; + match expected { + ExpectedRead::Missing => assert!(result.unwrap().is_none()), + ExpectedRead::Malformed => assert!(matches!(result, Err(Error::MalformedPayload))), + ExpectedRead::NonString => assert!(matches!(result, Err(Error::NonStringValue))), + } +} + +#[rstest] +#[tokio::test] +async fn base_manager_context_overrides_vault_location_and_data_key( + token_values: Vec<(&str, &str)>, +) { + let server: MockServer = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/alternate/data/managed/name")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({"api_token": "value"}))), + ) + .expect(2) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/v1/alternate/data/managed/name")) + .and(body_json(json!({ + "data": {"api_token": "updated", "description": "Managed key"} + }))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "data": { + "created_time": "", + "deletion_time": "", + "custom_metadata": null, + "destroyed": false, + "version": 2 + }, + "lease_id": "", + "lease_duration": 0, + "renewable": false, + "request_id": "", + "warnings": null, + "wrap_info": null + }))) + .expect(1) + .mount(&server) + .await; + Mock::given(method("DELETE")) + .and(path("/v1/alternate/data/managed/name")) + .respond_with(ResponseTemplate::new(204)) + .expect(1) + .mount(&server) + .await; + let manager: HashicorpVault = manager(&server, &token_values); + let operation = HashicorpOperationContext { + mount: Some(" /alternate/ ".to_owned()), + path_prefix: Some(" /managed/ ".to_owned()), + data_key: Some("api_token".to_owned()), + ..HashicorpOperationContext::default() + }; + let write_context = SecretWriteContext { + description: Some("Managed key".to_owned()), + operation: operation.clone(), + ..SecretWriteContext::default() + }; + + assert_eq!( + BaseSecretManager::async_read_secret(&manager, "name", &operation) + .await + .unwrap() + .unwrap() + .expose(), + "value" + ); + SecretWriter::async_write_secret( + &manager, + "name", + &SecretValue::new("updated"), + &write_context, + ) + .await + .unwrap(); + assert!( + BaseSecretManager::async_read_secret(&manager, "name", &operation) + .await + .unwrap() + .is_some() + ); + SecretDeleter::async_delete_secret(&manager, "name", &operation) + .await + .unwrap(); +} + +#[rstest] +#[tokio::test] +async fn rejects_description_that_would_replace_the_secret_value(token_values: Vec<(&str, &str)>) { + let server: MockServer = MockServer::start().await; + let manager: HashicorpVault = manager(&server, &token_values); + let context = SecretWriteContext { + description: Some("metadata".to_owned()), + operation: HashicorpOperationContext { + data_key: Some("description".to_owned()), + ..HashicorpOperationContext::default() + }, + ..SecretWriteContext::default() + }; + + assert!(matches!( + manager + .async_write_secret_with_context("name", &SecretValue::new("secret"), &context) + .await, + Err(Error::DataKeyConflictsWithDescription) + )); + assert!(server.received_requests().await.unwrap().is_empty()); +} + +#[rstest] +#[tokio::test] +async fn reads_cache_each_data_key_for_the_same_vault_path(token_values: Vec<(&str, &str)>) { + let server: MockServer = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/name")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({ + "key": "primary", + "alternate": "secondary" + }))), + ) + .expect(2) + .mount(&server) + .await; + let manager: HashicorpVault = manager(&server, &token_values); + let alternate = HashicorpOperationContext { + data_key: Some("alternate".to_owned()), + ..HashicorpOperationContext::default() + }; + + assert_eq!( + manager + .async_read_secret("name") + .await + .unwrap() + .unwrap() + .expose(), + "primary" + ); + assert_eq!( + BaseSecretManager::async_read_secret(&manager, "name", &alternate) + .await + .unwrap() + .unwrap() + .expose(), + "secondary" + ); + assert_eq!( + manager + .async_read_secret("name") + .await + .unwrap() + .unwrap() + .expose(), + "primary" + ); +} + +#[rstest] +#[tokio::test] +async fn base_manager_context_timeout_limits_vault_io(token_values: Vec<(&str, &str)>) { + let server: MockServer = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/name")) + .respond_with( + ResponseTemplate::new(200) + .set_delay(Duration::from_millis(100)) + .set_body_json(read_response(json!({"key": "value"}))), + ) + .expect(1) + .mount(&server) + .await; + let manager: HashicorpVault = manager(&server, &token_values); + let context = HashicorpOperationContext { + timeout: Some(Duration::from_millis(10)), + ..HashicorpOperationContext::default() + }; + + assert!(matches!( + BaseSecretManager::async_read_secret(&manager, "name", &context).await, + Err(Error::Timeout) + )); +} + +#[rstest] +#[tokio::test] +async fn concurrent_reads_share_a_load_but_verification_fetches_fresh( + token_values: Vec<(&str, &str)>, +) { + use litellm_secrets_types::SecretRotator; + use std::sync::atomic::{AtomicUsize, Ordering}; + let server = MockServer::start().await; + let reads = AtomicUsize::new(0); + Mock::given(method("GET")) + .and(path("/v1/secret/data/name")) + .respond_with(move |_: &wiremock::Request| { + let value = if reads.fetch_add(1, Ordering::SeqCst) == 0 { + "old" + } else { + "new" + }; + ResponseTemplate::new(200) + .set_body_json(read_response(json!({"key": value}))) + .set_delay(Duration::from_millis(20)) + }) + .expect(2) + .mount(&server) + .await; + let manager = manager(&server, &token_values); + let (first, second) = tokio::join!( + manager.async_read_secret("name"), + manager.async_read_secret("name") + ); + assert_eq!(first.unwrap().unwrap().expose(), "old"); + assert_eq!(second.unwrap().unwrap().expose(), "old"); + assert_eq!( + manager + .async_read_secret("name") + .await + .unwrap() + .unwrap() + .expose(), + "old" + ); + assert_eq!( + manager + .async_read_secret_fresh("name", &HashicorpOperationContext::default()) + .await + .unwrap() + .unwrap() + .expose(), + "new" + ); + assert_eq!( + manager + .async_read_secret("name") + .await + .unwrap() + .unwrap() + .expose(), + "new" + ); +} diff --git a/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/support.rs b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/support.rs new file mode 100644 index 00000000000..30e46f93248 --- /dev/null +++ b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/support.rs @@ -0,0 +1,175 @@ +use super::*; + +pub(super) fn config(server: &MockServer, values: &[(&str, &str)]) -> HashicorpVaultConfig { + let mut environment_values: HashMap = values + .iter() + .map(|(name, value)| ((*name).to_owned(), (*value).to_owned())) + .collect(); + environment_values.insert("HCP_VAULT_ADDR".to_owned(), server.uri()); + let environment: Arc = + Arc::new(move |name: &str| environment_values.get(name).cloned()); + HashicorpVaultConfig::from_environment(environment.as_ref()).unwrap() +} + +pub(super) fn manager(server: &MockServer, values: &[(&str, &str)]) -> HashicorpVault { + HashicorpVault::from_config(config(server, values), true).unwrap() +} + +pub(super) fn auth_response(token: &str, lease_duration: u64) -> serde_json::Value { + json!({ + "auth": { + "client_token": token, + "accessor": "", + "policies": [], + "token_policies": [], + "metadata": null, + "lease_duration": lease_duration, + "renewable": false, + "entity_id": "", + "token_type": "service", + "orphan": false + }, + "lease_id": "", + "lease_duration": lease_duration, + "renewable": false, + "request_id": "", + "warnings": null, + "wrap_info": null + }) +} + +pub(super) fn read_response(data: serde_json::Value) -> serde_json::Value { + json!({ + "data": { + "data": data, + "metadata": { + "created_time": "", + "deletion_time": "", + "custom_metadata": null, + "destroyed": false, + "version": 1 + } + }, + "lease_id": "", + "lease_duration": 0, + "renewable": false, + "request_id": "", + "warnings": null, + "wrap_info": null + }) +} + +pub(super) fn metadata_response(version: u64) -> serde_json::Value { + json!({ + "data": { + "cas_required": true, + "created_time": "", + "current_version": version, + "delete_version_after": "0s", + "max_versions": 0, + "oldest_version": 1, + "updated_time": "", + "custom_metadata": null, + "versions": {} + }, + "lease_id": "", + "lease_duration": 0, + "renewable": false, + "request_id": "", + "warnings": null, + "wrap_info": null + }) +} + +pub(super) fn write_response(version: u64) -> serde_json::Value { + json!({ + "data": { + "created_time": "", + "deletion_time": "", + "custom_metadata": null, + "destroyed": false, + "version": version + }, + "lease_id": "", + "lease_duration": 0, + "renewable": false, + "request_id": "", + "warnings": null, + "wrap_info": null + }) +} + +#[fixture] +pub(super) fn token_values() -> Vec<(&'static str, &'static str)> { + vec![("HCP_VAULT_TOKEN", "token")] +} + +#[derive(Deserialize)] +pub(super) struct ParityCase { + pub(super) env: HashMap, + pub(super) expected_secret_url: String, + pub(super) expected_login_url: Option, + pub(super) expected_login_namespace: Option, + pub(super) expected_secret_namespace: Option, + pub(super) secret_name: String, +} + +#[fixture] +pub(super) fn parity_cases() -> Vec { + serde_json::from_str(include_str!(concat!( + env!("CARGO_MANIFEST_DIR"), + "/../../../tests/test_litellm/secret_managers/hashicorp_vault_parity.json" + ))) + .unwrap() +} + +pub(super) const TEST_CERTIFICATE: &str = "-----BEGIN CERTIFICATE----- +MIIDDzCCAfegAwIBAgIUeMzLFLM/mRbPGbNAew5N2UTscocwDQYJKoZIhvcNAQEL +BQAwFzEVMBMGA1UEAwwMbGl0ZWxsbS10ZXN0MB4XDTI2MDkyMTIwMjA1OVoXDTI2 +MDkyMjIwMjA1OVowFzEVMBMGA1UEAwwMbGl0ZWxsbS10ZXN0MIIBIjANBgkqhkiG +9w0BAQEFAAOCAQ8AMIIBCgKCAQEAveYoSUJXybmkHmQsBfhBcv2Ob5Oy8ejZu+B3 +vTnrPumW4ANi1XXKBSazRGB3fEtAgr+3KhKeHaSKEQeBwJkAEBfdmQv0tpXICwHs +1kFNtU0owy54HVW5/ia+LMszsFcPzVIoMnbUOuiKr9RaV7P+IEFzILPBVuV4DoYH +yocjD3+9QNqokWgNL8LK37JijmNEFVaKFz0X6SyL2VRDlfPWTEBK52Gp/pvDgA6G +eTSfyI+kCm9h5ECTYUAtmatk9WPVS8sWOqV1EXVanFyYBU+mDxoywAS1/6CHeIPh +bNmCOZjPoO9qWBJ7ZyGhOconBigXY8qnlXymev+44IPHrx4urwIDAQABo1MwUTAd +BgNVHQ4EFgQUvaZrZ6HKtbr3ekeZmgy4b5Pq95QwHwYDVR0jBBgwFoAUvaZrZ6HK +tbr3ekeZmgy4b5Pq95QwDwYDVR0TAQH/BAUwAwEB/zANBgkqhkiG9w0BAQsFAAOC +AQEAEejrD8d1qDxW55XxQ4IC31rufoEvDV955jyvh2kALPaN/i5oWsBGI+UAQZna +aaoQXwzlmHrtDUBWl0LztVTUamIleUep2+PLLauqqt43vxppxMX8Jn2mnPO20YE/ +hIzGx0jN/LBG8PDyLSvHdlgjP9ofA4Vg4rTQugdXRgOvlCE/epnH/MADcg9KYJtJ +C1RObCIkL3LcdUbjStJRCY/U/FeWcgyncEPz95OFDkbrlNDajb6o6CkYfouqvhTc +8XlgjjAVKIbAbRgbVu3elsquuFM97x2DzWDjkrMNmDt1FJ9ubK36gL6B3o0UMaoQ +00R7x/eqvH+EkWa/2ekW9lpleQ== +-----END CERTIFICATE----- +"; + +pub(super) const TEST_PRIVATE_KEY: &str = "-----BEGIN PRIVATE KEY----- +MIIEvQIBADANBgkqhkiG9w0BAQEFAASCBKcwggSjAgEAAoIBAQC95ihJQlfJuaQe +ZCwF+EFy/Y5vk7Lx6Nm74He9Oes+6ZbgA2LVdcoFJrNEYHd8S0CCv7cqEp4dpIoR +B4HAmQAQF92ZC/S2lcgLAezWQU21TSjDLngdVbn+Jr4syzOwVw/NUigydtQ66Iqv +1FpXs/4gQXMgs8FW5XgOhgfKhyMPf71A2qiRaA0vwsrfsmKOY0QVVooXPRfpLIvZ +VEOV89ZMQErnYan+m8OADoZ5NJ/Ij6QKb2HkQJNhQC2Zq2T1Y9VLyxY6pXURdVqc +XJgFT6YPGjLABLX/oId4g+Fs2YI5mM+g72pYEntnIaE5yicGKBdjyqeVfKZ6/7jg +g8evHi6vAgMBAAECggEAGdJjlP6b8Fa5bdaCM/ebcrbuuNZVJVbb0JPHxGfNSLs7 +pE9hj5QaOdQW2Uviw3h6F61ZCzQH4xD+Iy2po5ZKb2XHYKnDB1bboj+LRGER337T +9aJqe9at2VTMVEv3Rdm40NsEk0QcPLxlK16NQFK90gYEUSSQPDAswJDSG2R/zHn+ +vADI907mW/goEJHeLn8PWGlNlSiR6x+5JJtq+GXCzUzVvJYQSCLGxCSl2x2H+0g7 +NhFI0zPpdzNmO/h+yhzaFb6Rp5U8+ZsnZ3qYjQ/03gw1myTDKJt1YaO9JvArnNYX +hcJQQ8Rt0bHhcrZA16bBOpqZlo5pKCicwI/netgN8QKBgQDcFz7AzdJ26sMSV32V +rwrMgIoggt8qDjO1ARwqW35A1TIge0FoW4M4KpsXQGGfT341uU1esXEcyZ/1L/5X +3ql2gX4DbOYLZLWYzZGR2hq33oi8HkhN98QrEwL9emSH8NqYX3Xxja3PrmCrSYJe +Zbnd9TIm2XkxyMoyXJu6M/QvnwKBgQDc4dzqTbxoGEGa5MuJoGmMwPnqgdG9UM5J +eExVnh7osxc2sOdsiPeRjjQTxs9v2kJwctC359OJoo9yGaaJeSghU4LEWJo1sqnA +fzSCLammYvtVAtniyNv5Mxk/6Uimi4NNDKaAKB+m4K2uSn3U9AmY7KPYMGaSbS9W +XSnobjxm8QKBgC8bPpAvvWs8ZhIn7bY659nLbUT2HeO3dHO6UBf0yzn/J6JyHxbB +93zvCZDZc8uQTRgcmCW7XtVlhjoJUqvl+Wlm39zF0xr/LCsPXKfWAb/2/lcdOCaP +8Emz4QD10EyUTYUtcWYJB/mafhBLRH8F0Nlj4J8WDu2L51MOJTqeYhZLAoGAWffN +icocAbJPlo22sdoa4+/+W5yBF8GAJMDRJtZ+9H1t6SLpQHYRkMIBSETkXUTjZvX9 +Ocs9iIQkNW9pO/mTdO+VBfCo71JUfknR02xR+6m5gYjlws/ZeYlssXGN2/hbhNiw +QOcW7Vv6olFJK6Iy/oz0t6wPO3kpnN3Zogi0paECgYEAwo44M1DdYCtV0snhmYM9 +5u0mPfYt5P2SVLXyUbr+vFTfrTL/WKnXIJgbsnj3Gvf+GIZv9tKcXhSNmEHQCYX4 +X3w9iTPddCHuvZ1fpufi2TyArJh0OkoNtLXJHTKrHjf2N+61AQzFiv5WieJrdE+H +qr32PTUuVGPyO9LyTY4/RL0= +-----END PRIVATE KEY----- +"; diff --git a/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/writes.rs b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/writes.rs new file mode 100644 index 00000000000..e2468e90c20 --- /dev/null +++ b/litellm-rust/crates/secrets-hashicorp/tests/secret_manager/writes.rs @@ -0,0 +1,542 @@ +use super::*; + +#[rstest] +#[tokio::test] +async fn write_and_delete_invalidate_the_read_cache(token_values: Vec<(&str, &str)>) { + use std::sync::atomic::{AtomicUsize, Ordering}; + let server = MockServer::start().await; + let revision = Arc::new(AtomicUsize::new(0)); + let current = revision.clone(); + Mock::given(method("GET")) + .and(path("/v1/secret/data/name")) + .respond_with( + move |_: &wiremock::Request| match current.load(Ordering::SeqCst) { + 0 => ResponseTemplate::new(200).set_body_json(read_response( + json!({"key": "old", "alternate": "old-alternate"}), + )), + 1 => ResponseTemplate::new(200).set_body_json(read_response( + json!({"key": "updated", "alternate": "updated-alternate"}), + )), + _ => ResponseTemplate::new(404).set_body_json(json!({"errors": ["missing"]})), + }, + ) + .expect(5) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/unrelated")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "unrelated"}))), + ) + .expect(1) + .mount(&server) + .await; + let written = revision.clone(); + Mock::given(method("POST")) + .and(path("/v1/secret/data/name")) + .respond_with(move |_: &wiremock::Request| { + written.store(1, Ordering::SeqCst); + ResponseTemplate::new(200).set_body_json(write_response(2)) + }) + .expect(1) + .mount(&server) + .await; + Mock::given(method("DELETE")) + .and(path("/v1/secret/data/name")) + .respond_with(move |_: &wiremock::Request| { + revision.store(2, Ordering::SeqCst); + ResponseTemplate::new(204) + }) + .expect(1) + .mount(&server) + .await; + let manager = manager(&server, &token_values); + let alternate = HashicorpOperationContext { + data_key: Some("alternate".into()), + ..Default::default() + }; + assert_eq!( + manager + .async_read_secret("name") + .await + .unwrap() + .unwrap() + .expose(), + "old" + ); + assert_eq!( + BaseSecretManager::async_read_secret(&manager, "name", &alternate) + .await + .unwrap() + .unwrap() + .expose(), + "old-alternate" + ); + assert_eq!( + manager + .async_read_secret("unrelated") + .await + .unwrap() + .unwrap() + .expose(), + "unrelated" + ); + manager + .async_write_secret("name", SecretValue::new("updated"), None) + .await + .unwrap(); + assert_eq!( + manager + .async_read_secret("name") + .await + .unwrap() + .unwrap() + .expose(), + "updated" + ); + assert_eq!( + BaseSecretManager::async_read_secret(&manager, "name", &alternate) + .await + .unwrap() + .unwrap() + .expose(), + "updated-alternate" + ); + manager.async_delete_secret("name").await.unwrap(); + assert!(manager.async_read_secret("name").await.unwrap().is_none()); + assert_eq!( + manager + .async_read_secret("unrelated") + .await + .unwrap() + .unwrap() + .expose(), + "unrelated" + ); +} + +#[rstest] +#[case::existing(200, 2)] +#[case::new_secret(404, 0)] +#[tokio::test] +async fn cas_required_writes_retry_with_the_current_version( + token_values: Vec<(&str, &str)>, + #[case] metadata_status: u16, + #[case] expected_cas: u64, +) { + let server: MockServer = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/secret/data/name")) + .and(body_json(json!({"data": {"key": "value"}}))) + .respond_with(ResponseTemplate::new(400).set_body_json(json!({"errors": ["CAS required"]}))) + .expect(1) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/v1/secret/metadata/name")) + .respond_with( + ResponseTemplate::new(metadata_status).set_body_json(metadata_response(expected_cas)), + ) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/v1/secret/data/name")) + .and(body_json(json!({ + "data": {"key": "value"}, + "options": {"cas": expected_cas} + }))) + .respond_with(ResponseTemplate::new(200).set_body_json(write_response(expected_cas + 1))) + .expect(1) + .mount(&server) + .await; + + let result = manager(&server, &token_values) + .async_write_secret("name", SecretValue::new("value"), None) + .await; + + assert!(result.is_ok()); +} + +#[rstest] +#[tokio::test] +async fn failed_cas_lookup_preserves_the_write_error(token_values: Vec<(&str, &str)>) { + let server: MockServer = MockServer::start().await; + Mock::given(method("POST")) + .and(path("/v1/secret/data/name")) + .respond_with( + ResponseTemplate::new(400).set_body_json(json!({"errors": ["write rejected"]})), + ) + .expect(1) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/v1/secret/metadata/name")) + .respond_with(ResponseTemplate::new(403).set_body_json(json!({"errors": ["forbidden"]}))) + .expect(1) + .mount(&server) + .await; + + let result = manager(&server, &token_values) + .async_write_secret("name", SecretValue::new("value"), None) + .await; + + assert!(matches!(result, Err(Error::Status { status: 400 }))); +} + +#[rstest] +#[tokio::test] +async fn rotation_applies_timeout_to_each_request(token_values: Vec<(&str, &str)>) { + let server = MockServer::start().await; + let timeout = Duration::from_secs(1); + let delay = timeout / 2; + Mock::given(method("GET")) + .and(header("X-Vault-Namespace", "team")) + .and(path("/v1/alternate/data/managed/current")) + .respond_with( + ResponseTemplate::new(200) + .set_delay(delay) + .set_body_json(read_response(json!({"api_token": "original"}))), + ) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(header("X-Vault-Namespace", "team")) + .and(path("/v1/alternate/data/managed/new")) + .and(body_json(json!({ + "data": {"api_token": "replacement", "description": "Rotated from current"} + }))) + .respond_with( + ResponseTemplate::new(200) + .set_delay(delay) + .set_body_json(json!({ + "data": { + "created_time": "", + "deletion_time": "", + "custom_metadata": null, + "destroyed": false, + "version": 1 + }, + "lease_id": "", + "lease_duration": 0, + "renewable": false, + "request_id": "", + "warnings": null, + "wrap_info": null + })), + ) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(header("X-Vault-Namespace", "team")) + .and(path("/v1/alternate/data/managed/new")) + .respond_with( + ResponseTemplate::new(200) + .set_delay(delay) + .set_body_json(read_response(json!({"api_token": "replacement"}))), + ) + .mount(&server) + .await; + Mock::given(method("DELETE")) + .and(header("X-Vault-Namespace", "team")) + .and(path("/v1/alternate/data/managed/current")) + .respond_with(ResponseTemplate::new(204).set_delay(delay)) + .mount(&server) + .await; + let manager = manager(&server, &token_values); + let context = HashicorpOperationContext { + namespace: Some("team".into()), + timeout: Some(timeout), + mount: Some("alternate".to_owned()), + path_prefix: Some("managed".to_owned()), + data_key: Some("api_token".to_owned()), + }; + + manager + .async_rotate_secret_with_context( + "current", + "new", + &SecretValue::new("replacement"), + &context, + ) + .await + .unwrap(); + let requests = server.received_requests().await.unwrap(); + let operations: Vec<_> = requests + .iter() + .map(|request| (request.method.as_str(), request.url.path())) + .collect(); + assert_eq!( + operations, + [ + ("GET", "/v1/alternate/data/managed/current"), + ("POST", "/v1/alternate/data/managed/new"), + ("GET", "/v1/alternate/data/managed/new"), + ("DELETE", "/v1/alternate/data/managed/current"), + ] + ); +} + +#[rstest] +#[tokio::test] +#[ignore] +async fn live_vault_round_trip() { + let environment: Arc = + Arc::new(litellm_core_utils::settings::ProcessEnvironment); + let manager: HashicorpVault = HashicorpVault::new(environment, true).unwrap(); + let name: String = std::env::var("LITELLM_VAULT_LIVE_SECRET_NAME").unwrap(); + let value: SecretValue = SecretValue::new("native-live-value"); + let location = manager.secret_location(&name).unwrap(); + println!( + "native provenance: {} vaultrs {} {:?} {} {}", + module_path!(), + manager.config().address, + location.namespace, + location.mount, + location.path + ); + manager + .async_write_secret(&name, value.clone(), None) + .await + .unwrap(); + assert_eq!( + manager.async_read_secret(&name).await.unwrap().unwrap(), + value + ); + let replacement = SecretValue::new("replacement-Ï€\n"); + manager + .async_write_secret(&name, replacement.clone(), None) + .await + .unwrap(); + assert_eq!( + manager.async_read_secret(&name).await.unwrap().unwrap(), + replacement + ); + manager.async_delete_secret(&name).await.unwrap(); + assert!(manager.async_read_secret(&name).await.unwrap().is_none()); +} + +#[rstest] +#[tokio::test] +async fn same_name_rotation_keeps_the_replacement(token_values: Vec<(&str, &str)>) { + use std::sync::atomic::{AtomicUsize, Ordering}; + let server = MockServer::start().await; + let reads = AtomicUsize::new(0); + Mock::given(method("GET")) + .and(path("/v1/secret/data/name")) + .respond_with(move |_: &wiremock::Request| { + let value = if reads.fetch_add(1, Ordering::SeqCst) == 0 { + "original" + } else { + "replacement" + }; + ResponseTemplate::new(200).set_body_json(read_response(json!({"key": value}))) + }) + .expect(2) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/v1/secret/data/name")) + .and(body_json(json!({"data": {"key": "replacement", "description": "Rotated from name"}}))) + .respond_with(ResponseTemplate::new(200).set_body_json(json!({ + "data": {"created_time": "", "deletion_time": "", "custom_metadata": null, "destroyed": false, "version": 2}, + "lease_id": "", "lease_duration": 0, "renewable": false, "request_id": "", "warnings": null, "wrap_info": null + }))).expect(1).mount(&server).await; + Mock::given(method("DELETE")) + .respond_with(ResponseTemplate::new(204)) + .expect(0) + .mount(&server) + .await; + let manager = manager(&server, &token_values); + manager + .async_rotate_secret("name", "name", &SecretValue::new("replacement")) + .await + .unwrap(); + assert_eq!( + manager + .async_read_secret("name") + .await + .unwrap() + .unwrap() + .expose(), + "replacement" + ); +} + +#[rstest] +#[case::verification(false)] +#[case::retirement(true)] +#[tokio::test] +async fn rotation_reports_partial_completion_without_losing_the_write_response( + token_values: Vec<(&str, &str)>, + #[case] verified: bool, +) { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/old")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({"key": "old"}))), + ) + .expect(1) + .mount(&server) + .await; + Mock::given(method("POST")) + .and(path("/v1/secret/data/new")) + .respond_with(ResponseTemplate::new(200).set_body_json(write_response(2))) + .expect(1) + .mount(&server) + .await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/new")) + .respond_with(ResponseTemplate::new(200).set_body_json(read_response( + json!({"key": if verified { "replacement" } else { "stale" }}), + ))) + .expect(1) + .mount(&server) + .await; + Mock::given(method("DELETE")) + .and(path("/v1/secret/data/old")) + .respond_with(ResponseTemplate::new(403).set_body_json(json!({"errors": ["denied"]}))) + .expect(u64::from(verified)) + .mount(&server) + .await; + let result = manager(&server, &token_values) + .async_rotate_secret("old", "new", &SecretValue::new("replacement")) + .await; + match result { + Err(RotationError::Verification { response, source }) if !verified => { + assert_eq!(response["version"], 2); + assert!(matches!( + source, + Error::Operation(litellm_secrets_types::Error::NewSecretMismatch) + )); + } + Err(RotationError::Retirement { response, source }) if verified => { + assert_eq!(response["version"], 2); + assert!(matches!(source, Error::Status { status: 403 })); + } + other => panic!("unexpected rotation outcome: {other:?}"), + } +} + +#[rstest] +#[case::different_namespace( + " /team-b/ ", + " /alternate/ ", + " /prefix/ ", + Some("team-b"), + "/v1/alternate/data/prefix/key" +)] +#[case::clear_namespace("", "", "", None, "/v1/secret/data/key")] +#[tokio::test] +async fn operation_overrides_isolate_cached_targets_and_apply_to_writes_and_deletes( + #[case] namespace: &str, + #[case] mount: &str, + #[case] prefix: &str, + #[case] expected_namespace: Option<&str>, + #[case] expected_path: &'static str, +) { + let server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/configured/data/configured/key")) + .and(header("X-Vault-Namespace", "team-a")) + .respond_with( + ResponseTemplate::new(200).set_body_json(read_response(json!({"key":"default-value"}))), + ) + .expect(1) + .mount(&server) + .await; + let namespace_header = expected_namespace.map(str::to_owned); + Mock::given(path(expected_path)) + .respond_with(move |request: &wiremock::Request| { + assert_eq!( + request + .headers + .get("X-Vault-Namespace") + .map(|value| value.to_str().unwrap()), + namespace_header.as_deref() + ); + match request.method.as_str() { + "GET" => ResponseTemplate::new(200) + .set_body_json(read_response(json!({"password":"override-value"}))), + "POST" => { + assert_eq!( + request.body_json::().unwrap(), + json!({"data":{"password":"written"}}) + ); + ResponseTemplate::new(200).set_body_json(write_response(2)) + } + "DELETE" => ResponseTemplate::new(204), + _ => panic!("unexpected method"), + } + }) + .expect(4) + .mount(&server) + .await; + let manager = manager( + &server, + &[ + ("HCP_VAULT_TOKEN", "token"), + ("HCP_VAULT_SECRET_NAMESPACE", "team-a"), + ("HCP_VAULT_MOUNT_NAME", "configured"), + ("HCP_VAULT_PATH_PREFIX", "configured"), + ], + ); + let context = HashicorpOperationContext { + namespace: Some(namespace.into()), + mount: Some(mount.into()), + path_prefix: Some(prefix.into()), + data_key: Some("password".into()), + ..Default::default() + }; + for _ in 0..2 { + assert_eq!( + manager + .async_read_secret("key") + .await + .unwrap() + .unwrap() + .expose(), + "default-value" + ); + assert_eq!( + BaseSecretManager::async_read_secret(&manager, "key", &context) + .await + .unwrap() + .unwrap() + .expose(), + "override-value" + ); + } + SecretWriter::async_write_secret( + &manager, + "key", + &SecretValue::new("written"), + &SecretWriteContext { + operation: context.clone(), + ..Default::default() + }, + ) + .await + .unwrap(); + SecretDeleter::async_delete_secret(&manager, "key", &context) + .await + .unwrap(); + assert_eq!( + BaseSecretManager::async_read_secret(&manager, "key", &context) + .await + .unwrap() + .unwrap() + .expose(), + "override-value" + ); + assert_eq!( + manager + .async_read_secret("key") + .await + .unwrap() + .unwrap() + .expose(), + "default-value" + ); +} diff --git a/litellm-rust/crates/secrets-types/Cargo.toml b/litellm-rust/crates/secrets-types/Cargo.toml index acd29746722..dcd06d1a741 100644 --- a/litellm-rust/crates/secrets-types/Cargo.toml +++ b/litellm-rust/crates/secrets-types/Cargo.toml @@ -7,6 +7,8 @@ repository.workspace = true [dependencies] litellm-auth-types.workspace = true +moka.workspace = true +tokio = { workspace = true, features = ["sync"] } serde.workspace = true serde_json.workspace = true thiserror.workspace = true diff --git a/litellm-rust/crates/secrets-types/src/base_secret_manager.rs b/litellm-rust/crates/secrets-types/src/base_secret_manager.rs index e8aed9f280e..737fad28662 100644 --- a/litellm-rust/crates/secrets-types/src/base_secret_manager.rs +++ b/litellm-rust/crates/secrets-types/src/base_secret_manager.rs @@ -1,4 +1,4 @@ -use crate::{Error, SecretOperationContext, SecretValue, SecretWriteContext}; +use crate::{Error, SecretValue, SecretWriteContext}; pub fn validate_secret_name(name: &str) -> Result<(), Error> { if name.split('/').any(|segment| segment == "..") @@ -11,64 +11,114 @@ pub fn validate_secret_name(name: &str) -> Result<(), Error> { Ok(()) } -#[expect( - async_fn_in_trait, - reason = "closed backend dispatch does not require Send bounds on generic rotation" -)] pub trait BaseSecretManager { type Error: From; - type WriteResponse; - type DeleteResponse; + type Context: Clone + Default + Send + Sync; - async fn async_read_secret( + fn async_read_secret( &self, name: &str, - context: &SecretOperationContext, - ) -> Result, Self::Error>; - async fn async_write_secret( + context: &Self::Context, + ) -> impl std::future::Future, Self::Error>> + Send; +} + +pub trait SecretWriter: BaseSecretManager { + type WriteResponse; + + fn async_write_secret( &self, name: &str, value: &SecretValue, - context: &SecretWriteContext, - ) -> Result; - async fn async_delete_secret( - &self, - name: &str, - recovery_window_in_days: Option, - context: &SecretOperationContext, - ) -> Result; + context: &SecretWriteContext, + ) -> impl std::future::Future> + Send; } -pub async fn async_rotate_secret( +pub trait SecretDeleter: BaseSecretManager { + type DeleteResponse; + + fn async_delete_secret( + &self, + name: &str, + context: &Self::Context, + ) -> impl std::future::Future> + Send; +} + +pub trait SecretRotator: SecretDeleter { + type RotationResponse; + + fn async_write_replacement( + &self, + current_name: &str, + new_name: &str, + value: &SecretValue, + context: &Self::Context, + ) -> impl std::future::Future> + Send; + + fn async_read_secret_fresh( + &self, + name: &str, + context: &Self::Context, + ) -> impl std::future::Future, Self::Error>> + Send; +} + +#[derive(Debug, PartialEq, Eq, thiserror::Error)] +pub enum RotationError { + #[error("could not read the current secret")] + Read(#[source] E), + #[error("replacement write failed; provider state may be unknown")] + Write(#[source] E), + #[error("replacement was written but could not be verified")] + Verification { + response: Box, + #[source] + source: E, + }, + #[error("replacement was verified but retiring the old secret failed")] + Retirement { + response: Box, + #[source] + source: E, + }, +} + +pub async fn async_rotate_secret( manager: &M, current_name: &str, new_name: &str, value: &SecretValue, - context: &SecretOperationContext, -) -> Result { + context: &M::Context, +) -> Result> { if manager - .async_read_secret(current_name, context) - .await? + .async_read_secret_fresh(current_name, context) + .await + .map_err(RotationError::Read)? .is_none() { - return Err(Error::CurrentSecretMissing.into()); + return Err(RotationError::Read(Error::CurrentSecretMissing.into())); } let response = manager - .async_write_secret( - new_name, - value, - &SecretWriteContext::rotated_from(current_name, context.clone()), - ) - .await?; - if manager - .async_read_secret(new_name, context) - .await? - .is_none() - { - return Err(Error::NewSecretMissing.into()); + .async_write_replacement(current_name, new_name, value, context) + .await + .map_err(RotationError::Write)?; + let verification = match manager.async_read_secret_fresh(new_name, context).await { + Ok(None) => Err(Error::NewSecretMissing.into()), + Ok(Some(actual)) if actual != *value => Err(Error::NewSecretMismatch.into()), + Ok(Some(_)) => Ok(()), + Err(error) => Err(error), + }; + if let Err(source) = verification { + return Err(RotationError::Verification { + response: Box::new(response), + source, + }); + } + if current_name != new_name + && let Err(source) = manager.async_delete_secret(current_name, context).await + { + return Err(RotationError::Retirement { + response: Box::new(response), + source, + }); } - manager - .async_delete_secret(current_name, Some(7), context) - .await?; Ok(response) } diff --git a/litellm-rust/crates/secrets-types/src/cache.rs b/litellm-rust/crates/secrets-types/src/cache.rs new file mode 100644 index 00000000000..524de1ed2f1 --- /dev/null +++ b/litellm-rust/crates/secrets-types/src/cache.rs @@ -0,0 +1,73 @@ +use std::{future::Future, hash::Hash, sync::Arc, time::Duration}; + +use moka::future::Cache; +use tokio::sync::Mutex; + +#[derive(Clone)] +pub struct SecretCache { + entries: Cache>>>, +} + +impl SecretCache +where + K: Eq + Hash + Clone + Send + Sync + 'static, + V: Clone + Send + Sync + 'static, +{ + pub fn new(capacity: u64, ttl: Duration) -> Self { + Self { + entries: Cache::builder() + .max_capacity(capacity) + .time_to_live(ttl) + .support_invalidation_closures() + .build(), + } + } + + pub async fn read( + &self, + key: K, + load: impl Future, E>>, + ) -> Result, E> { + let entry = self + .entries + .get_with(key, async { Arc::new(Mutex::new(None)) }) + .await; + let mut value = entry.lock().await; + if value.is_some() { + return Ok(value.clone()); + } + // Invalidated loads only populate their detached entry, never the cache's replacement. + let loaded = load.await?; + *value = loaded.clone(); + Ok(loaded) + } + + pub async fn invalidate(&self, key: &K) { + self.entries.invalidate(key).await; + } + + pub async fn refresh( + &self, + key: K, + load: impl Future, E>>, + ) -> Result, E> { + let entry = Arc::new(Mutex::new(None)); + let mut value = entry.lock().await; + self.entries.insert(key, entry.clone()).await; + let loaded = load.await?; + *value = loaded.clone(); + Ok(loaded) + } + + pub async fn insert(&self, key: K, value: V) { + self.entries + .insert(key, Arc::new(Mutex::new(Some(value)))) + .await; + } + + pub fn invalidate_where(&self, predicate: impl Fn(&K) -> bool + Send + Sync + 'static) { + self.entries + .invalidate_entries_if(move |key, _| predicate(key)) + .expect("invalidation closures are enabled"); + } +} diff --git a/litellm-rust/crates/secrets-types/src/context.rs b/litellm-rust/crates/secrets-types/src/context.rs index 126815ab32b..1ec493edf42 100644 --- a/litellm-rust/crates/secrets-types/src/context.rs +++ b/litellm-rust/crates/secrets-types/src/context.rs @@ -1,21 +1,41 @@ use std::{collections::BTreeMap, time::Duration}; -use crate::SecretValue; +use crate::{Error, KeyManagementSystem, SecretValue}; #[derive(Clone, Debug, Default, Eq, PartialEq)] pub enum SecretOperationContext { #[default] Default, Aws(AwsOperationContext), + Azure(AzureOperationContext), + Google(GoogleOperationContext), Hashicorp(HashicorpOperationContext), Cyberark(CyberarkOperationContext), } impl SecretOperationContext { + pub fn validate_for(&self, system: KeyManagementSystem) -> Result<(), Error> { + let compatible = match self { + Self::Default => true, + Self::Aws(_) => system == KeyManagementSystem::AwsSecretManager, + Self::Azure(_) => system == KeyManagementSystem::AzureKeyVault, + Self::Google(_) => system == KeyManagementSystem::GoogleSecretManager, + Self::Hashicorp(_) => system == KeyManagementSystem::HashicorpVault, + Self::Cyberark(_) => system == KeyManagementSystem::Cyberark, + }; + if compatible { + Ok(()) + } else { + Err(Error::InvalidOperationContext) + } + } + pub fn timeout(&self) -> Option { match self { Self::Default => None, Self::Aws(context) => context.timeout, + Self::Azure(context) => context.timeout, + Self::Google(context) => context.timeout, Self::Hashicorp(context) => context.timeout, Self::Cyberark(context) => context.timeout, } @@ -24,6 +44,9 @@ impl SecretOperationContext { #[derive(Clone, Debug, Default, Eq, PartialEq)] pub struct AwsOperationContext { + pub access_key_id: Option, + pub secret_access_key: Option, + pub session_token: Option, pub timeout: Option, pub region_name: Option, pub role_name: Option, @@ -32,10 +55,12 @@ pub struct AwsOperationContext { pub profile_name: Option, pub web_identity_token: Option, pub sts_endpoint: Option, + pub bedrock_runtime_endpoint: Option, } #[derive(Clone, Debug, Default, Eq, PartialEq)] pub struct HashicorpOperationContext { + pub namespace: Option, pub timeout: Option, pub mount: Option, pub path_prefix: Option, @@ -48,14 +73,14 @@ pub struct CyberarkOperationContext { } #[derive(Clone, Debug, Default, Eq, PartialEq)] -pub struct SecretWriteContext { +pub struct SecretWriteContext { pub description: Option, pub tags: BTreeMap, - pub operation: SecretOperationContext, + pub operation: C, } -impl SecretWriteContext { - pub fn rotated_from(current_name: &str, operation: SecretOperationContext) -> Self { +impl SecretWriteContext { + pub fn rotated_from(current_name: &str, operation: C) -> Self { Self { description: Some(format!("Rotated from {current_name}")), tags: BTreeMap::new(), @@ -63,3 +88,13 @@ impl SecretWriteContext { } } } + +#[derive(Clone, Debug, Default, Eq, PartialEq)] +pub struct AzureOperationContext { + pub timeout: Option, +} + +#[derive(Clone, Debug, Default, Eq, PartialEq)] +pub struct GoogleOperationContext { + pub timeout: Option, +} diff --git a/litellm-rust/crates/secrets-types/src/error.rs b/litellm-rust/crates/secrets-types/src/error.rs index cae9c7f4c69..027dd54d8cc 100644 --- a/litellm-rust/crates/secrets-types/src/error.rs +++ b/litellm-rust/crates/secrets-types/src/error.rs @@ -6,4 +6,8 @@ pub enum Error { CurrentSecretMissing, #[error("new secret could not be verified")] NewSecretMissing, + #[error("new secret does not match the replacement")] + NewSecretMismatch, + #[error("secret manager received an incompatible operation context")] + InvalidOperationContext, } diff --git a/litellm-rust/crates/secrets-types/src/lib.rs b/litellm-rust/crates/secrets-types/src/lib.rs index d330fc7bbea..4d7eaf8313e 100644 --- a/litellm-rust/crates/secrets-types/src/lib.rs +++ b/litellm-rust/crates/secrets-types/src/lib.rs @@ -1,17 +1,22 @@ #![forbid(unsafe_code)] mod base_secret_manager; +mod cache; mod config; mod context; mod error; mod value; -pub use base_secret_manager::{BaseSecretManager, async_rotate_secret, validate_secret_name}; +pub use base_secret_manager::{ + BaseSecretManager, RotationError, SecretDeleter, SecretRotator, SecretWriter, + async_rotate_secret, validate_secret_name, +}; +pub use cache::SecretCache; pub use config::{AccessMode, KeyManagementSettings, KeyManagementSystem}; pub use context::{ - AwsOperationContext, CyberarkOperationContext, HashicorpOperationContext, - SecretOperationContext, SecretWriteContext, + AwsOperationContext, AzureOperationContext, CyberarkOperationContext, GoogleOperationContext, + HashicorpOperationContext, SecretOperationContext, SecretWriteContext, }; pub use error::Error; pub use litellm_auth_types::SecretValue; -pub use value::Secret; +pub use value::{PythonSecretRead, Secret}; diff --git a/litellm-rust/crates/secrets-types/src/value.rs b/litellm-rust/crates/secrets-types/src/value.rs index 087537fb3eb..11f57ef4448 100644 --- a/litellm-rust/crates/secrets-types/src/value.rs +++ b/litellm-rust/crates/secrets-types/src/value.rs @@ -7,6 +7,12 @@ pub enum Secret { Json(#[redact] serde_json::Value), } +#[derive(Debug)] +pub enum PythonSecretRead { + Value(Option), + PrimaryJson(SecretValue), +} + impl From for Secret { fn from(value: SecretValue) -> Self { Self::String(value) diff --git a/litellm-rust/crates/secrets-types/tests/cache.rs b/litellm-rust/crates/secrets-types/tests/cache.rs new file mode 100644 index 00000000000..cec4f5198e3 --- /dev/null +++ b/litellm-rust/crates/secrets-types/tests/cache.rs @@ -0,0 +1,162 @@ +use std::{convert::Infallible, time::Duration}; + +use litellm_secrets_types::SecretCache; +use rstest::rstest; +use tokio::sync::oneshot; + +#[rstest] +#[case::delete(false)] +#[case::write(true)] +#[tokio::test] +async fn an_old_load_cannot_restore_a_mutated_entry(#[case] write: bool) { + let cache = SecretCache::new(10, Duration::from_secs(60)); + let (started, loading) = oneshot::channel(); + let (release, finish) = oneshot::channel(); + let old = cache.read("key", async { + started.send(()).unwrap(); + finish.await.unwrap(); + Ok::<_, Infallible>(Some("old")) + }); + let mutation = async { + loading.await.unwrap(); + if write { + cache.insert("key", "new").await; + } else { + cache.invalidate(&"key").await; + } + release.send(()).unwrap(); + }; + let (old_result, ()) = tokio::join!(old, mutation); + assert_eq!(old_result.unwrap(), Some("old")); + let current = cache + .read("key", async { Ok::<_, Infallible>(None) }) + .await + .unwrap(); + assert_eq!(current, write.then_some("new")); +} + +#[tokio::test] +async fn location_invalidation_detaches_all_projections_and_keeps_other_secrets() { + let cache = SecretCache::new(10, Duration::from_secs(60)); + cache.insert(("target", "first"), "old").await; + cache.insert(("unrelated", "first"), "retained").await; + let (started, loading) = oneshot::channel(); + let (release, finish) = oneshot::channel(); + let old = cache.read(("target", "second"), async { + started.send(()).unwrap(); + finish.await.unwrap(); + Ok::<_, Infallible>(Some("old")) + }); + let mutation = async { + loading.await.unwrap(); + cache.invalidate_where(|(location, _)| *location == "target"); + release.send(()).unwrap(); + }; + let (result, ()) = tokio::join!(old, mutation); + assert_eq!(result.unwrap(), Some("old")); + for projection in ["first", "second"] { + assert_eq!( + cache + .read(("target", projection), async { Ok::<_, Infallible>(None) }) + .await + .unwrap(), + None + ); + } + assert_eq!( + cache + .read(("unrelated", "first"), async { Ok::<_, Infallible>(None) }) + .await + .unwrap(), + Some("retained") + ); +} + +#[tokio::test] +async fn concurrent_misses_share_a_load_and_cancellation_allows_a_retry() { + let cache = SecretCache::new(10, Duration::from_secs(60)); + let (started, loading) = oneshot::channel(); + let (release, finish) = oneshot::channel(); + let first = cache.read("key", async { + started.send(()).unwrap(); + finish.await.unwrap(); + Ok::<_, Infallible>(Some("value")) + }); + let second = async { + loading.await.unwrap(); + release.send(()).unwrap(); + cache.read("key", async { panic!("duplicate load") }).await + }; + let (first, second): (_, Result<_, Infallible>) = tokio::join!(first, second); + assert_eq!(first.unwrap(), Some("value")); + assert_eq!(second.unwrap(), Some("value")); + + let (started, loading) = oneshot::channel(); + let cancelled = cache.read("cancelled", async { + started.send(()).unwrap(); + std::future::pending::, Infallible>>().await + }); + tokio::select! { + _ = loading => {}, + _ = cancelled => panic!("load must remain pending"), + } + assert_eq!( + cache + .read("cancelled", async { Ok::<_, Infallible>(Some("retry")) }) + .await + .unwrap(), + Some("retry") + ); +} + +#[tokio::test] +async fn refresh_bypasses_cached_values_and_errors_and_absence_are_retried() { + let cache = SecretCache::new(10, Duration::from_secs(60)); + cache.insert("key", "old").await; + assert_eq!( + cache + .refresh("key", async { Ok::<_, Infallible>(Some("fresh")) }) + .await + .unwrap(), + Some("fresh") + ); + assert_eq!( + cache + .read("key", async { Ok::<_, Infallible>(None) }) + .await + .unwrap(), + Some("fresh") + ); + for result in [Err("failure"), Ok(None), Ok(Some("recovered"))] { + assert_eq!(cache.read("retry", async { result }).await, result); + } +} + +#[tokio::test] +async fn expired_entries_reload_and_empty_values_are_cached() { + let cache = SecretCache::new(10, Duration::from_secs(60)); + assert_eq!( + cache + .read("empty", async { Ok::<_, Infallible>(Some("")) }) + .await + .unwrap(), + Some("") + ); + assert_eq!( + cache + .read("empty", async { Ok::<_, Infallible>(Some("changed")) }) + .await + .unwrap(), + Some("") + ); + let expiring = SecretCache::new(10, Duration::from_nanos(1)); + expiring.insert("key", "old").await; + tokio::time::sleep(Duration::from_millis(1)).await; + assert_eq!( + expiring + .read("key", async { Ok::<_, Infallible>(Some("new")) }) + .await + .unwrap(), + Some("new") + ); +} diff --git a/litellm-rust/crates/secrets-types/tests/context.rs b/litellm-rust/crates/secrets-types/tests/context.rs index 7c9e91b3a9e..a167e87954f 100644 --- a/litellm-rust/crates/secrets-types/tests/context.rs +++ b/litellm-rust/crates/secrets-types/tests/context.rs @@ -1,8 +1,9 @@ use std::{collections::BTreeMap, time::Duration}; use litellm_secrets_types::{ - AwsOperationContext, CyberarkOperationContext, HashicorpOperationContext, - SecretOperationContext, SecretValue, SecretWriteContext, + AwsOperationContext, AzureOperationContext, CyberarkOperationContext, GoogleOperationContext, + HashicorpOperationContext, KeyManagementSystem, SecretOperationContext, SecretValue, + SecretWriteContext, }; use rstest::{fixture, rstest}; @@ -108,3 +109,40 @@ fn rotation_write_context_preserves_the_operation_context(aws_context: SecretOpe assert!(context.tags.is_empty()); assert_eq!(context.operation, aws_context); } + +#[rstest] +#[case( + KeyManagementSystem::AwsSecretManager, + SecretOperationContext::Aws(Default::default()) +)] +#[case(KeyManagementSystem::AzureKeyVault, SecretOperationContext::Azure(AzureOperationContext { timeout: Some(Duration::from_secs(1)) }))] +#[case(KeyManagementSystem::GoogleSecretManager, SecretOperationContext::Google(GoogleOperationContext { timeout: Some(Duration::from_secs(1)) }))] +#[case( + KeyManagementSystem::HashicorpVault, + SecretOperationContext::Hashicorp(Default::default()) +)] +#[case( + KeyManagementSystem::Cyberark, + SecretOperationContext::Cyberark(Default::default()) +)] +fn provider_context_accepts_only_its_owner( + #[case] owner: KeyManagementSystem, + #[case] context: SecretOperationContext, +) { + for system in [ + KeyManagementSystem::AwsSecretManager, + KeyManagementSystem::AzureKeyVault, + KeyManagementSystem::GoogleSecretManager, + KeyManagementSystem::HashicorpVault, + KeyManagementSystem::Cyberark, + ] { + assert_eq!(context.validate_for(system).is_ok(), system == owner); + assert!(SecretOperationContext::Default.validate_for(system).is_ok()); + } + if matches!( + owner, + KeyManagementSystem::AzureKeyVault | KeyManagementSystem::GoogleSecretManager + ) { + assert_eq!(context.timeout(), Some(Duration::from_secs(1))); + } +} diff --git a/litellm-rust/crates/secrets-types/tests/rotation.rs b/litellm-rust/crates/secrets-types/tests/rotation.rs index 1bc5e0a9eeb..36c1a065ffa 100644 --- a/litellm-rust/crates/secrets-types/tests/rotation.rs +++ b/litellm-rust/crates/secrets-types/tests/rotation.rs @@ -1,23 +1,53 @@ use std::sync::atomic::{AtomicUsize, Ordering}; use litellm_secrets_types::{ - BaseSecretManager, Error, HashicorpOperationContext, SecretOperationContext, SecretValue, - SecretWriteContext, async_rotate_secret, validate_secret_name, + BaseSecretManager, Error, HashicorpOperationContext, RotationError, SecretDeleter, + SecretOperationContext, SecretRotator, SecretValue, SecretWriteContext, SecretWriter, + async_rotate_secret, validate_secret_name, }; use rstest::{fixture, rstest}; struct Manager { step: AtomicUsize, absent_at: Option, + verified_value: &'static str, operation: SecretOperationContext, + delete_error: bool, + fail_at: Option, } impl BaseSecretManager for Manager { type Error = Error; - type WriteResponse = &'static str; - type DeleteResponse = (); + type Context = SecretOperationContext; async fn async_read_secret( + &self, + _name: &str, + _context: &Self::Context, + ) -> Result, Error> { + panic!("rotation must bypass cached reads") + } +} + +impl SecretRotator for Manager { + type RotationResponse = &'static str; + + async fn async_write_replacement( + &self, + current_name: &str, + new_name: &str, + value: &SecretValue, + context: &Self::Context, + ) -> Result { + self.async_write_secret( + new_name, + value, + &SecretWriteContext::rotated_from(current_name, context.clone()), + ) + .await + } + + async fn async_read_secret_fresh( &self, name: &str, context: &SecretOperationContext, @@ -25,8 +55,21 @@ impl BaseSecretManager for Manager { assert_eq!(context, &self.operation); let step = self.step.fetch_add(1, Ordering::SeqCst); assert_eq!(name, if step == 0 { "old" } else { "new" }); - Ok((self.absent_at != Some(step)).then(|| SecretValue::new("value"))) + if self.fail_at == Some(step) { + return Err(Error::UnsafeSecretName); + } + Ok((self.absent_at != Some(step)).then(|| { + SecretValue::new(if step == 0 { + "value" + } else { + self.verified_value + }) + })) } +} + +impl SecretWriter for Manager { + type WriteResponse = &'static str; async fn async_write_secret( &self, @@ -35,6 +78,9 @@ impl BaseSecretManager for Manager { context: &SecretWriteContext, ) -> Result { assert_eq!(self.step.fetch_add(1, Ordering::SeqCst), 1); + if self.fail_at == Some(1) { + return Err(Error::UnsafeSecretName); + } assert_eq!(name, "new"); assert_eq!(value.expose(), "replacement"); assert_eq!(context.description.as_deref(), Some("Rotated from old")); @@ -42,18 +88,24 @@ impl BaseSecretManager for Manager { assert_eq!(context.operation, self.operation); Ok("provider-response") } +} + +impl SecretDeleter for Manager { + type DeleteResponse = (); async fn async_delete_secret( &self, name: &str, - recovery_window_in_days: Option, context: &SecretOperationContext, ) -> Result<(), Error> { assert_eq!(self.step.fetch_add(1, Ordering::SeqCst), 3); assert_eq!(name, "old"); - assert_eq!(recovery_window_in_days, Some(7)); assert_eq!(context, &self.operation); - Ok(()) + if self.delete_error { + Err(Error::UnsafeSecretName) + } else { + Ok(()) + } } } @@ -75,7 +127,10 @@ async fn rotation_verifies_before_deleting_and_returns_provider_response( ) { let manager = Manager { step: AtomicUsize::new(0), + delete_error: false, + fail_at: None, absent_at: None, + verified_value: "replacement", operation: operation.clone(), }; assert_eq!( @@ -88,18 +143,21 @@ async fn rotation_verifies_before_deleting_and_returns_provider_response( } #[rstest] -#[case::current_secret_missing(0, Error::CurrentSecretMissing, 1)] -#[case::new_secret_missing(2, Error::NewSecretMissing, 3)] +#[case::current_secret_missing(0, RotationError::Read(Error::CurrentSecretMissing), 1)] +#[case::new_secret_missing(2, RotationError::Verification { response: Box::new("provider-response"), source: Error::NewSecretMissing }, 3)] #[tokio::test] async fn missing_old_or_new_value_stops_rotation_before_deletion( replacement: SecretValue, #[case] absent_at: usize, - #[case] expected: Error, + #[case] expected: RotationError<&'static str, Error>, #[case] calls: usize, ) { let manager = Manager { step: AtomicUsize::new(0), + delete_error: false, + fail_at: None, absent_at: Some(absent_at), + verified_value: "replacement", operation: SecretOperationContext::default(), }; assert_eq!( @@ -156,3 +214,88 @@ fn names_reject_path_traversal_and_control_characters(#[case] name: &str) { fn names_allow_safe_values(#[case] name: &str) { assert_eq!(validate_secret_name(name), Ok(())); } + +#[tokio::test] +async fn a_different_replacement_never_deletes_the_current_secret() { + let manager = Manager { + step: AtomicUsize::new(0), + delete_error: false, + fail_at: None, + absent_at: None, + verified_value: "stale-value", + operation: SecretOperationContext::Default, + }; + assert_eq!( + async_rotate_secret( + &manager, + "old", + "new", + &SecretValue::new("replacement"), + &SecretOperationContext::Default + ) + .await, + Err(RotationError::Verification { + response: Box::new("provider-response"), + source: Error::NewSecretMismatch + }) + ); + assert_eq!(manager.step.load(Ordering::SeqCst), 3); +} + +#[tokio::test] +async fn failed_retirement_preserves_the_verified_replacement_response() { + let manager = Manager { + step: AtomicUsize::new(0), + absent_at: None, + verified_value: "replacement", + operation: SecretOperationContext::Default, + delete_error: true, + fail_at: None, + }; + assert_eq!( + async_rotate_secret( + &manager, + "old", + "new", + &SecretValue::new("replacement"), + &SecretOperationContext::Default + ) + .await, + Err(RotationError::Retirement { + response: Box::new("provider-response"), + source: Error::UnsafeSecretName + }) + ); +} + +#[rstest] +#[case::read(0, RotationError::Read(Error::UnsafeSecretName))] +#[case::write(1, RotationError::Write(Error::UnsafeSecretName))] +#[case::verification(2, RotationError::Verification { response: Box::new("provider-response"), source: Error::UnsafeSecretName })] +#[tokio::test] +async fn provider_failures_stop_rotation_before_retirement( + #[case] fail_at: usize, + #[case] expected: RotationError<&'static str, Error>, +) { + let manager = Manager { + step: AtomicUsize::new(0), + absent_at: None, + verified_value: "replacement", + operation: SecretOperationContext::default(), + delete_error: false, + fail_at: Some(fail_at), + }; + assert_eq!( + async_rotate_secret( + &manager, + "old", + "new", + &SecretValue::new("replacement"), + &SecretOperationContext::default() + ) + .await + .unwrap_err(), + expected + ); + assert_eq!(manager.step.load(Ordering::SeqCst), fail_at + 1); +} diff --git a/litellm-rust/crates/secrets/Cargo.toml b/litellm-rust/crates/secrets/Cargo.toml index e5f30025976..17acc01682b 100644 --- a/litellm-rust/crates/secrets/Cargo.toml +++ b/litellm-rust/crates/secrets/Cargo.toml @@ -14,6 +14,8 @@ azure = ["dep:litellm-secrets-azure"] cyberark = ["dep:litellm-secrets-cyberark"] [dependencies] +futures-util.workspace = true +litellm-python-compat = { path = "../python-compat" } litellm-secrets-types.workspace = true litellm-secrets-aws = { workspace = true, optional = true } litellm-secrets-google = { workspace = true, optional = true } @@ -24,7 +26,6 @@ litellm-core-utils.workspace = true base64.workspace = true serde.workspace = true strum.workspace = true -jsonwebtoken.workspace = true serde_json.workspace = true thiserror.workspace = true reqwest.workspace = true diff --git a/litellm-rust/crates/secrets/PARITY.md b/litellm-rust/crates/secrets/PARITY.md new file mode 100644 index 00000000000..aeed4ba4b83 --- /dev/null +++ b/litellm-rust/crates/secrets/PARITY.md @@ -0,0 +1,277 @@ +# Python secret-manager test parity + +The inventory covers 132 tests in the secret-manager suites and the legacy secret-manager utility suites. It maps 108 tests to Rust coverage and identifies 24 tests owned by other layers or live environments. Rust tests exercise requests, returned values, caching, routing and failure behavior. Multiple Python tests can map to one parameterized Rust test + +Tests of Python extension lifecycle, SDK credential selection in `auth-azure`, proxy hooks and example subclasses remain at their owning boundary. They are called out below rather than counted as Rust secret-manager coverage. Default AWS partition endpoints are owned by the AWS SDK + +Azure callback absence remains `None`, while a native HTTP 404 still permits environment fallback. `azure_callback_absence_preserves_none_but_errors_fall_back` and `python_read_failures_preserve_provider_fallback_rules` cover these distinct results + +Native APIs preserve typed errors and explicit absence. Python-compatible reads restore Python fallback, coercion, Google negative caching and missing-value behavior. Vault namespaces use the SDK namespace header instead of Python’s equivalent URL prefix. Rotation verifies fresh provider reads instead of trusting a just-written cache entry, preventing deletion after a failed replacement + +Google payload corruption is intentionally rejected: malformed base64 and mismatched CRC32C values fail without populating the cache. `failed_or_missing_reads_are_not_cached` tests this correction against Python's permissive decoder and omitted checksum validation. See [RFC 4648 section 3.3](https://www.rfc-editor.org/rfc/rfc4648#section-3.3) and [Google's integrity guidance](https://docs.cloud.google.com/secret-manager/docs/data-integrity) + +The Python API audit found that earlier AWS fallback tests encoded the wrong expectation. Differential calls to the existing handler show that missing secrets, denied reads, missing string payloads, and missing or empty primary secrets return `None`; they do not activate environment fallback or defaults. `test_aws_absence_and_failed_reads_match_python_without_environment_fallback` compares the public getter under both dispatch decisions, including HTTP 500 responses and their exact request counts. Python-compatible AWS reads disable SDK retries, matching Python's single attempt for service errors while retaining the native API's retry configuration. The bridge resolver test `aws_read_failure_preserves_absence_without_environment_fallback` checks the same rule when environment values exist. `test_aws_primary_values_match_python_handler` checks typed values, including arbitrary-size integers. `test_aws_primary_json_errors_preserve_python_exception_details` checks the exception class, arguments, document, and position against Python + +Public AWS, Vault, CyberArk, and Google read methods now use catalog selection. AWS, Vault, and CyberArk keep their Python coroutine entrypoints, while Rust receives provider operation contexts. Full API replacement remains incomplete: AWS write/delete/rotation methods still need native bindings. Vault and CyberArk mutations now use native dispatch and share their native read caches. SDK-client configuration capture also needs to preserve explicit credentials, endpoints, and regions. Timeout phase handling, AWS optional-parameter side effects, per-call region selection without a base region, and environment lookup timing need further parity work. The private binding returns a Future, while the public methods retain ordinary Python coroutines as verified by lazy execution and `asyncio.create_task` tests. This follows the separation described in the PyO3 [signature](https://pyo3.rs/v0.29.2/function/signature.html) and [async](https://pyo3.rs/v0.29.2/async-await.html) guides. The catalog remains Python-only while these gaps are open + +## Public API audit + +The rollout decision comes from [catalog.py](../../../litellm/rust_bridge/catalog.py). All secret-manager rules remain `PYTHON_ONLY`; `LITELLM_RUST` does not make these incomplete routes production-ready. Differential tests pass explicit rules into the dispatch boundary + +| Python entrypoint | Native bridge coverage | Remaining API work | +| --- | --- | --- | +| `litellm.get_secret`, `get_secret_str`, `get_secret_bool` | Existing Python entrypoints dispatch supported manager reads | SDK-client configuration, environment read timing and complete failure conversion | +| AWS `sync_read_secret`, `async_read_secret`, primary-secret helpers | Public signatures and coroutines retained; credentials, absence and typed JSON tested | Option consumption, timeout phases and region resolution | +| AWS `async_write_secret`, `async_delete_secret`, `async_rotate_secret`, `async_replicate_secret`, `async_put_secret_value` | Rust provider operations exist | Public native dispatch, original response fields and Python error contracts | +| Vault `sync_read_secret`, `async_read_secret` | Public signatures, nested overrides, namespace and data-key cache isolation tested | Complete timeout and initialization error parity | +| Vault `async_write_secret`, `async_delete_secret`, `async_rotate_secret` | Public native dispatch, complete response envelopes, HTTP error dictionaries, timeouts and fresh verification tested | Authentication and input-conversion edge cases, transport retries, timeout phases and mutable configuration read points | +| CyberArk reads, writes, deletes and rotations | Public native dispatch, shared cache, coroutine behavior, status errors and request counts tested | Other transport failures and client initialization timing | +| Google `get_secret_from_google_secret_manager` | Public native dispatch and distinct initial/cached missing results | Credential configuration and complete error parity | +| Azure Key Vault, AWS KMS, Google KMS SDK clients | Global secret-handler dispatch supports recognized clients | Explicit SDK credentials, endpoints, regions and caller-supplied credentials | +| Custom managers and subclasses | Preserve Python callbacks | Caller implementations must never be replaced by built-in native managers | + +`_SecretManagerRuntime` is a private implementation detail, not a replacement SDK class. Its async methods return Futures; public `async def` methods retain lazy coroutine creation and `asyncio.create_task` support. Passing the same names and arguments is insufficient to claim parity until the remaining return-value, error, cache and configuration differences above are closed + +## [tests/test_litellm/secret_managers/test_aws_secret_manager_replication.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_replication.py) + +| Python test | Rust coverage or boundary | +| --- | --- | +| `test_write_secret_replicates_when_configured` | [creation_replicates_only_to_configured_regions](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_write_secret_no_replication_when_not_configured` | [creation_replicates_only_to_configured_regions](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_replication_failure_does_not_fail_write` | [creation_passes_tags_and_kms_and_survives_replication_failure](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_async_replicate_secret_empty_regions_returns_empty` | [creation_passes_tags_and_kms_and_survives_replication_failure](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_async_replicate_secret_correct_payload` | [direct_replication_returns_response_or_service_error](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_replication_fires_on_create` | [creation_replicates_only_to_configured_regions](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_load_aws_secret_manager_passes_replica_regions` | [creation_replicates_only_to_configured_regions](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_write_secret_http_error_raises` | [create_failure_does_not_overwrite_an_alias_without_a_deletion_date](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_write_secret_timeout_raises` | [write_and_replication_timeouts_remain_errors](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_replicate_secret_http_error_raises` | [direct_replication_returns_response_or_service_error](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_replicate_secret_timeout_raises` | [write_and_replication_timeouts_remain_errors](../secrets-aws/tests/secret_manager/writes.rs) | + +## [tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_rotation.py) + +| Python test | Rust coverage or boundary | +| --- | --- | +| `test_rotate_secret_same_name_writes_requested_value_in_place` | [same_name_rotation_uses_put_and_returns_its_response](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_rotate_secret_different_names_persists_requested_value_and_deletes_old_alias` | [renamed_rotation_reads_creates_verifies_then_deletes](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_rotate_secret_back_to_name_inside_recovery_window_restores_and_stores_new_value` | [recovery_window_alias_is_restored_updated_and_tagged](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_write_secret_to_name_inside_recovery_window_reschedules_deletion_when_update_fails` | [failed_update_reschedules_deletion_of_a_restored_alias](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_write_secret_to_name_inside_recovery_window_restores_and_stores_new_value` | [recovery_window_alias_is_restored_updated_and_tagged](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_write_secret_to_live_existing_name_still_fails_without_overwriting` | [create_failure_does_not_overwrite_an_alias_without_a_deletion_date](../secrets-aws/tests/secret_manager/writes.rs) | + +## [tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py](../../../tests/test_litellm/secret_managers/test_aws_secret_manager_v2.py) + +| Python test | Rust coverage or boundary | +| --- | --- | +| `test_create_secret_uses_customer_managed_kms_key_from_settings` | [creation_replicates_only_to_configured_regions](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_create_secret_omits_kms_key_id_when_not_configured` | [write_read_delete_preserves_the_complete_secret_string](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_write_and_read_json_secret` | [write_read_delete_preserves_the_complete_secret_string](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_prepare_request_builds_partition_endpoint` | AWS SDK owns partition endpoint construction. LiteLLM region selection is exercised by `trait_read_uses_the_aws_region_from_its_operation_context`; no vendor endpoint table is duplicated | +| `test_prepare_request_explicit_bedrock_runtime_endpoint_param_still_wins` | [endpoint_overrides_replace_the_service_and_override_the_region](../secrets-aws/tests/secret_manager/configuration.rs) | +| `test_prepare_request_env_bedrock_runtime_endpoint_still_wins` | [endpoint_overrides_replace_the_service_and_override_the_region](../secrets-aws/tests/secret_manager/configuration.rs) | + +## [tests/test_litellm/secret_managers/test_base_secret_manager.py](../../../tests/test_litellm/secret_managers/test_base_secret_manager.py) + +| Python test | Rust coverage or boundary | +| --- | --- | +| `test_raise_if_unsafe_secret_name_rejects_traversal_and_line_breaks` | [names_reject_path_traversal_and_control_characters](../secrets-types/tests/rotation.rs) | +| `test_raise_if_unsafe_secret_name_allows_legitimate_aliases` | [names_allow_safe_values](../secrets-types/tests/rotation.rs) | + +## [tests/test_litellm/secret_managers/test_custom_secret_manager.py](../../../tests/test_litellm/secret_managers/test_custom_secret_manager.py) + +| Python test | Rust coverage or boundary | +| --- | --- | +| `test_custom_secret_manager_initialization` | Exercises the Python example subclass itself or Python default methods. Caller-authored Python implementations remain Python callbacks; resolver integration is covered by `manager_strings_are_coerced_like_literal_eval` | +| `test_custom_secret_manager_sync_read` | Exercises the Python example subclass itself or Python default methods. Caller-authored Python implementations remain Python callbacks; resolver integration is covered by `manager_strings_are_coerced_like_literal_eval` | +| `test_custom_secret_manager_async_read` | Exercises the Python example subclass itself or Python default methods. Caller-authored Python implementations remain Python callbacks; resolver integration is covered by `manager_strings_are_coerced_like_literal_eval` | +| `test_custom_secret_manager_async_write` | Exercises the Python example subclass itself or Python default methods. Caller-authored Python implementations remain Python callbacks; resolver integration is covered by `manager_strings_are_coerced_like_literal_eval` | +| `test_custom_secret_manager_async_delete` | Exercises the Python example subclass itself or Python default methods. Caller-authored Python implementations remain Python callbacks; resolver integration is covered by `manager_strings_are_coerced_like_literal_eval` | +| `test_custom_secret_manager_integration_with_litellm` | [manager_strings_are_coerced_like_literal_eval](../secrets/tests/resolution.rs) | +| `test_minimal_custom_secret_manager` | Exercises the Python example subclass itself or Python default methods. Caller-authored Python implementations remain Python callbacks; resolver integration is covered by `manager_strings_are_coerced_like_literal_eval` | + +## [tests/test_litellm/secret_managers/test_cyberark_secret_manager.py](../../../tests/test_litellm/secret_managers/test_cyberark_secret_manager.py) + +| Python test | Rust coverage or boundary | +| --- | --- | +| `test_sync_read_matches_parity_fixture` | [secret_names_use_python_quote_encoding](../secrets-cyberark/tests/secret_manager/reads.rs) | +| `test_async_write_matches_parity_fixture` | [writes_match_python_parity_fixture](../secrets-cyberark/tests/secret_manager/writes.rs) | +| `test_missing_credentials_raise_value_error` | [new_validates_credentials_before_license_and_configuration](../secrets-cyberark/tests/secret_manager/configuration.rs) | + +## [tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py](../../../tests/test_litellm/secret_managers/test_get_azure_ad_token_provider.py) + +| Python test | Rust coverage or boundary | +| --- | --- | +| `test_deployment_identity_reaches_workload_and_managed_identity_only` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate | +| `test_deployment_identity_survives_a_developer_only_token_credentials_setting` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate | +| `test_default_azure_credential_keeps_its_full_chain` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate | +| `test_deployment_identity_refuses_to_mint_a_token_for_a_configured_service_principal` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate | +| `test_deployment_identity_still_reaches_a_system_assigned_managed_identity` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate | +| `test_deployment_identity_keeps_the_user_assigned_identity_under_a_dev_only_setting` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate | +| `test_get_azure_ad_token_provider_client_secret_credential` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate | +| `test_get_azure_ad_token_provider_managed_identity_credential` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate | +| `test_get_azure_ad_token_provider_certificate_credential` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate | +| `test_get_azure_ad_token_provider_password_protected_certificate_credential` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate | +| `test_get_azure_ad_token_provider_default_azure_credential` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate | +| `test_get_azure_ad_token_provider_prefers_workload_identity_over_managed_identity` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate | +| `test_get_azure_ad_token_provider_defaults_to_default_azure_credential` | Credential-selection contract belongs to `litellm-auth-azure`, not a secrets crate | + +## [tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py](../../../tests/test_litellm/secret_managers/test_hashicorp_secret_manager.py) + +| Python test | Rust coverage or boundary | +| --- | --- | +| `test_sync_read_uses_login_namespace_for_approle_and_secret_namespace_for_url` | [login_and_secret_namespaces_follow_python_precedence](../secrets-hashicorp/tests/secret_manager/configuration.rs) | +| `test_login_header_is_omitted_when_no_namespace_is_configured` | [login_and_secret_namespaces_follow_python_precedence](../secrets-hashicorp/tests/secret_manager/configuration.rs) | +| `test_sync_read_per_secret_namespace_overrides_secret_namespace` | [operation_overrides_isolate_cached_targets_and_apply_to_writes_and_deletes](../secrets-hashicorp/tests/secret_manager/writes.rs) | +| `test_sync_read_caches_per_resolved_target` | [operation_overrides_isolate_cached_targets_and_apply_to_writes_and_deletes](../secrets-hashicorp/tests/secret_manager/writes.rs) | +| `test_sync_read_caches_per_data_key_for_the_same_secret_path` | [reads_cache_each_data_key_for_the_same_vault_path](../secrets-hashicorp/tests/secret_manager/reads.rs) | +| `test_async_delete_evicts_every_cached_field_of_the_secret_path` | [write_and_delete_invalidate_the_read_cache](../secrets-hashicorp/tests/secret_manager/writes.rs) | +| `test_async_read_uses_secret_namespace_and_login_namespace` | [login_and_secret_namespaces_follow_python_precedence](../secrets-hashicorp/tests/secret_manager/configuration.rs) | +| `test_async_write_and_read_share_the_secret_namespace_target` | [write_and_delete_invalidate_the_read_cache](../secrets-hashicorp/tests/secret_manager/writes.rs) | +| `test_tls_login_uses_login_namespace` | [tls_login_posts_the_role_and_uses_the_client_identity](../secrets-hashicorp/tests/secret_manager/configuration.rs) | +| `test_configuration_matches_native_parity_fixture` | [configuration_matches_python_parity_fixture](../secrets-hashicorp/tests/secret_manager/configuration.rs) | + +## [tests/test_litellm/secret_managers/test_secret_manager_handler.py](../../../tests/test_litellm/secret_managers/test_secret_manager_handler.py) + +| Python test | Rust coverage or boundary | +| --- | --- | +| `test_azure_key_vault_matches_rust_parity_fixture` | [parity_fixture_matches_python_backend_contract](../secrets-azure/tests/key_vault.rs) | + +## [tests/test_litellm/secret_managers/test_secret_managers_main.py](../../../tests/test_litellm/secret_managers/test_secret_managers_main.py) + +| Python test | Rust coverage or boundary | +| --- | --- | +| `test_oidc_google_success` | [google_expiry_caps_cache_and_preserves_audience](../secrets/tests/oidc.rs) | +| `test_oidc_google_cached` | [google_expiry_caps_cache_and_preserves_audience](../secrets/tests/oidc.rs) | +| `test_oidc_google_cache_ttl_capped_by_token_exp` | [google_tokens_expire_at_the_python_cache_deadline](../secrets/tests/oidc.rs) | +| `test_oidc_google_expired_token_not_cached` | [google_expiry_caps_cache_and_preserves_audience](../secrets/tests/oidc.rs) | +| `test_oidc_google_long_lived_token_still_capped_at_default_ttl` | [google_tokens_expire_at_the_python_cache_deadline](../secrets/tests/oidc.rs) | +| `test_oidc_google_non_jwt_token_keeps_default_ttl` | [google_tokens_expire_at_the_python_cache_deadline](../secrets/tests/oidc.rs) | +| `test_oidc_google_failure` | [google_oidc_failures_are_not_cached_or_hidden_by_defaults](../secrets/tests/oidc.rs) | +| `test_oidc_circleci_success` | [environment_sources_resolve_expected_value](../secrets/tests/oidc.rs) | +| `test_oidc_circleci_failure` | [missing_oidc_environment_is_an_error](../secrets/tests/oidc.rs) | +| `test_oidc_github_success` | [github_requests_are_authenticated_cached_and_revalidate_environment](../secrets/tests/oidc.rs) | +| `test_oidc_github_missing_env` | [github_requests_are_authenticated_cached_and_revalidate_environment](../secrets/tests/oidc.rs) | +| `test_oidc_azure_file_success` | [file_allowlist_resolves_symlinks_while_environment_paths_remain_explicit](../secrets/tests/oidc.rs) | +| `test_oidc_azure_ad_token_success` | [azure_oidc_acquires_the_requested_scope_and_preserves_failures](../secrets/tests/oidc.rs) | +| `test_oidc_file_success` | [file_allowlist_resolves_symlinks_while_environment_paths_remain_explicit](../secrets/tests/oidc.rs) | +| `test_oidc_file_rejects_path_outside_allowlist` | [file_allowlist_resolves_symlinks_while_environment_paths_remain_explicit](../secrets/tests/oidc.rs) | +| `test_oidc_file_rejects_relative_path` | [file_allowlist_resolves_symlinks_while_environment_paths_remain_explicit](../secrets/tests/oidc.rs) | +| `test_oidc_env_success` | [environment_sources_resolve_expected_value](../secrets/tests/oidc.rs) | +| `test_oidc_env_path_success` | [file_allowlist_resolves_symlinks_while_environment_paths_remain_explicit](../secrets/tests/oidc.rs) | +| `test_unsupported_oidc_provider` | [invalid_references_fail_before_environment_lookup](../secrets/tests/oidc.rs) | +| `test_normalize_nonempty_secret_str` | [normalization_matches_python_without_changing_embedded_whitespace](../secrets/tests/resolution.rs) | +| `test_secret_manager_would_be_consulted_matches_get_secret` | [gating_prediction_matches_actual_lookup](../secrets/tests/aws.rs) | +| `test_secret_manager_would_be_consulted_is_false_without_a_client` | [prefix_is_removed_once_and_resolved_from_environment](../secrets/tests/resolution.rs) | + +## [tests/litellm_utils_tests/test_secret_manager.py](../../../tests/litellm_utils_tests/test_secret_manager.py) + +| Python test | Rust coverage or boundary | +| --- | --- | +| `test_aws_secret_manager` | [write_read_delete_preserves_the_complete_secret_string](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_oidc_google` | [google_expiry_caps_cache_and_preserves_audience](../secrets/tests/oidc.rs) | +| `test_oidc_github` | [github_requests_are_authenticated_cached_and_revalidate_environment](../secrets/tests/oidc.rs) | +| `test_oidc_circleci` | [environment_sources_resolve_expected_value](../secrets/tests/oidc.rs) | +| `test_oidc_circleci_v2` | [environment_sources_resolve_expected_value](../secrets/tests/oidc.rs) | +| `test_oidc_circleci_with_azure` | Quarantined live Azure token exchange, outside secrets crates. CircleCI token retrieval is covered by `environment_sources_resolve_expected_value` | +| `test_oidc_circle_v1_with_amazon` | Quarantined live AWS token exchange, outside secrets crates. Token retrieval and STS forwarding are covered independently | +| `test_oidc_env_variable` | [environment_sources_resolve_expected_value](../secrets/tests/oidc.rs) | +| `test_oidc_file` | [file_allowlist_resolves_symlinks_while_environment_paths_remain_explicit](../secrets/tests/oidc.rs) | +| `test_oidc_env_path` | [file_allowlist_resolves_symlinks_while_environment_paths_remain_explicit](../secrets/tests/oidc.rs) | +| `test_google_secret_manager` | [successful_reads_use_auth_latest_version_and_cache_including_empty_values](../secrets-google/tests/secret_manager.rs) | +| `test_google_secret_manager_read_in_memory` | [python_reads_reuse_cached_absence_until_expiry](../secrets-google/tests/secret_manager.rs) | +| `test_should_read_secret_from_secret_manager` | [gating_prediction_matches_actual_lookup](../secrets/tests/aws.rs) | +| `test_get_secret_with_access_mode` | [gating_prediction_matches_actual_lookup](../secrets/tests/aws.rs) | +| `test_key_management_settings_defaults` | [config_preserves_defaults_nulls_and_serialized_names](../secrets-types/tests/config.rs) | +| `test_key_management_settings_custom_values` | [config_preserves_defaults_nulls_and_serialized_names](../secrets-types/tests/config.rs) | +| `test_async_write_secret_receives_description_and_tags` | Proxy hook behavior stays in Python. Native write metadata is covered by `trait_write_uses_typed_write_context` | +| `test_key_management_settings_serialization_roundtrip` | [config_preserves_defaults_nulls_and_serialized_names](../secrets-types/tests/config.rs) | + +## [tests/litellm_utils_tests/test_get_secret.py](../../../tests/litellm_utils_tests/test_get_secret.py) + +| Python test | Rust coverage or boundary | +| --- | --- | +| `test_azure_kms` | [azure_handler_reads_missing_and_failed_secrets](../secrets/tests/azure.rs) | + +## [tests/litellm_utils_tests/test_aws_secret_manager.py](../../../tests/litellm_utils_tests/test_aws_secret_manager.py) + +| Python test | Rust coverage or boundary | +| --- | --- | +| `test_write_and_read_simple_secret` | [write_read_delete_preserves_the_complete_secret_string](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_write_and_read_json_secret` | [write_read_delete_preserves_the_complete_secret_string](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_read_nonexistent_secret` | [failed_read_returns_none_but_invalid_primary_json_is_an_error](../secrets-aws/tests/secret_manager/reads.rs) | +| `test_primary_secret_functionality` | [primary_lookup_preserves_read_semantics](../secrets-aws/tests/secret_manager/reads.rs) | +| `test_write_secret_with_description_and_tags` | [creation_passes_tags_and_kms_and_survives_replication_failure](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_secret_manager_with_iam_role_settings` | [configured_sts_credentials_sign_the_secret_request](../secrets-aws/tests/secret_manager/configuration.rs) | +| `test_secret_manager_with_cross_account_settings` | [configured_sts_credentials_sign_the_secret_request](../secrets-aws/tests/secret_manager/configuration.rs) | +| `test_secret_manager_with_irsa_settings` | [configured_sts_credentials_sign_the_secret_request](../secrets-aws/tests/secret_manager/configuration.rs) | +| `test_secret_manager_with_custom_sts_endpoint` | [configured_sts_credentials_sign_the_secret_request](../secrets-aws/tests/secret_manager/configuration.rs) | +| `test_secret_manager_with_aws_profile` | [configured_profile_credentials_override_static_environment_credentials](../secrets-aws/tests/secret_manager/configuration.rs) | +| `test_load_aws_secret_manager_with_settings` | [creation_replicates_only_to_configured_regions](../secrets-aws/tests/secret_manager/writes.rs) | +| `test_end_to_end_iam_role_secret_write` | Live AWS account test, not a unit test. Offline STS signing and secret writes are covered without account assumptions | + +## [tests/litellm_utils_tests/test_hashicorp.py](../../../tests/litellm_utils_tests/test_hashicorp.py) + +| Python test | Rust coverage or boundary | +| --- | --- | +| `test_hashicorp_secret_manager_get_secret` | [token_reads_use_vault_headers_and_cache_values](../secrets-hashicorp/tests/secret_manager/reads.rs) | +| `test_hashicorp_secret_manager_write_secret` | [write_and_delete_invalidate_the_read_cache](../secrets-hashicorp/tests/secret_manager/writes.rs) | +| `test_hashicorp_secret_manager_write_secret_with_team_overrides` | [operation_overrides_isolate_cached_targets_and_apply_to_writes_and_deletes](../secrets-hashicorp/tests/secret_manager/writes.rs) | +| `test_hashicorp_secret_manager_delete_secret` | [write_and_delete_invalidate_the_read_cache](../secrets-hashicorp/tests/secret_manager/writes.rs) | +| `test_hashicorp_secret_manager_delete_secret_with_team_overrides` | [operation_overrides_isolate_cached_targets_and_apply_to_writes_and_deletes](../secrets-hashicorp/tests/secret_manager/writes.rs) | +| `test_hashicorp_secret_manager_tls_cert_auth` | [tls_login_posts_the_role_and_uses_the_client_identity](../secrets-hashicorp/tests/secret_manager/configuration.rs) | +| `test_hashicorp_secret_manager_approle_auth` | [approle_login_uses_namespace_and_reuses_the_token](../secrets-hashicorp/tests/secret_manager/configuration.rs) | +| `test_hashicorp_custom_mount_and_prefix` | [namespace_mount_and_prefix_are_sanitized_in_the_url](../secrets-hashicorp/tests/secret_manager/reads.rs) | +| `test_hashicorp_get_url_rejects_path_traversal` | [no_auth_and_invalid_names_fail_without_requests](../secrets-hashicorp/tests/secret_manager/configuration.rs) | +| `test_hashicorp_secret_manager_rotate_secret_different_names` | [rotation_applies_timeout_to_each_request](../secrets-hashicorp/tests/secret_manager/writes.rs) | +| `test_hashicorp_secret_manager_rotate_secret_same_name` | [same_name_rotation_keeps_the_replacement](../secrets-hashicorp/tests/secret_manager/writes.rs) | +| `test_hashicorp_secret_manager_rotate_secret_current_not_found` | [missing_old_or_new_value_stops_rotation_before_deletion](../secrets-types/tests/rotation.rs) | +| `test_hashicorp_secret_manager_rotate_secret_write_fails` | [provider_failures_stop_rotation_before_retirement](../secrets-types/tests/rotation.rs) | +| `test_hashicorp_secret_manager_rotate_secret_with_team_overrides` | [rotation_applies_timeout_to_each_request](../secrets-hashicorp/tests/secret_manager/writes.rs) | +| `test_hashicorp_secret_manager_rotate_secret_value_mismatch` | [a_different_replacement_never_deletes_the_current_secret](../secrets-types/tests/rotation.rs) | + +## [tests/litellm_utils_tests/test_cyberark.py](../../../tests/litellm_utils_tests/test_cyberark.py) + +| Python test | Rust coverage or boundary | +| --- | --- | +| `test_cyberark_write_secret_rejects_yaml_injection` | [unsafe_names_fail_before_http_calls](../secrets-cyberark/tests/secret_manager/reads.rs) | +| `test_cyberark_ensure_variable_exists_escapes_yaml_metacharacters` | [policy_writes_preserve_yaml_metacharacters_as_one_variable](../secrets-cyberark/tests/secret_manager/writes.rs) | +| `test_cyberark_write_and_read_secret` | [writes_tolerate_policy_status_and_cache_value](../secrets-cyberark/tests/secret_manager/writes.rs) | +| `test_cyberark_rotate_secret` | [rotation_stores_the_replacement_and_retains_other_aliases](../secrets-cyberark/tests/secret_manager/writes.rs) | +| `test_cyberark_rotate_secret_with_new_alias` | [rotation_stores_the_replacement_and_retains_other_aliases](../secrets-cyberark/tests/secret_manager/writes.rs) | + +## Public provider read boundary + +`test_public_aws_reads_preserve_coroutines_and_per_call_credentials` verifies lazy coroutine execution, `asyncio.create_task`, positional and keyword calls, request payloads, and per-call credentials, region, and endpoint overrides. `test_public_aws_primary_reads_ignore_operation_overrides_like_python` retains Python's ignored primary-read overrides. Bootstrap keys bypass only synchronous reads, including native backend initialization + +Real HTTP timeouts are swallowed by AWS reads because LiteLLM's standard HTTP handler raises `litellm.Timeout`; tests that inject `httpx.TimeoutException` bypass that wrapping. `test_public_aws_read_timeouts_follow_the_python_http_handler` compares both implementations against delayed responses + +Vault reads retain nested overrides, Python string conversion, and cache isolation. CyberArk reuses authentication and preserves raw secret text in its cache. Python's shared cache JSON-decodes CyberArk values on subsequent reads, corrupting quoted strings and changing types. The native behavior intentionally fixes this corruption, with the Python difference shown in `test_public_cyberark_reads_reuse_authentication_and_cached_values` + +Google's first missing-secret read raises, while a cached miss returns `None`. `python_reads_reuse_cached_absence_until_expiry` and `python_cached_absence_expires_and_allows_recovery` retain both outcomes + + +## Public CyberArk mutation boundary + +Public writes, deletes, and rotations retain the Python method signatures and coroutine entrypoints. Write and delete results preserve Python's status/message dictionaries. Unsupported deletion clears the shared native read cache without a provider request. Authentication and write HTTP failures preserve Python's messages and request counts, including the ignored initial authentication failure while ensuring a policy. Python-compatible reads and writes do not retry HTTP 401; native Rust retry policies remain unchanged + +The bridge compares connection-refused errors against Python and builds HTTP status messages through HTTPX. Other transport failures and client-initialization timing still need a complete API audit; these checks do not establish full error parity + +`test_public_cyberark_writes_and_deletes_share_the_read_cache` verifies real HTTP writes followed by sync and async cached reads and deletion invalidation. `test_public_cyberark_write_errors_match_python_without_http_retries`, `test_public_cyberark_connection_errors_match_python`, and `test_cyberark_handler_errors_match_python_after_cached_authentication_is_denied` compare error results against Python. Missing-extension selection remains covered for each mutation + +Rotation preserves the documented fresh-read safeguard. Python's base rotation checks a cache populated by the write and does not compare the stored value, so a successful response can conceal a missing or incorrect replacement. `test_public_cyberark_rotation_requires_a_fresh_matching_replacement` rejects both cases and verifies the old cache remains intact. `test_public_cyberark_rotation_stops_after_a_failed_write` preserves the old value and returns the write error without further requests. Conjur retains the old provider alias because its deletion API is unsupported + + +## Public Vault mutation boundary + +Public Vault writes, deletes and rotation now use catalog selection. Their Python signatures and coroutine entrypoints stay unchanged. Native operations use the Vault SDK request types and authenticated client while retaining the complete response bytes for Python JSON conversion. This preserves additional response fields, key order and arbitrary-size integers. Python-compatible writes make one provider attempt; the existing native API keeps its CAS recovery behavior + +`test_public_vault_writes_preserve_complete_responses_and_request_fields` checks the response envelope, nested operation settings, payload and ignored tags. `test_public_vault_mutation_http_errors_match_python_without_retry` compares write/delete error dictionaries, including namespace URLs and exact request counts. Rotation tests compare current-secret failures, ordered request paths, write failures, failed verification, mismatched values, malformed verification shapes, same-name updates and best-effort old-alias deletion. A successful HTTP response containing `status: error` stops rotation and returns the original response. Verification fields remain raw JSON until Python error conversion, preserving large integers, nested values and the distinction between integer and floating-point type errors. Unsupported extension selection is checked separately for write, delete and rotate + +`test_public_vault_mutation_timeouts_match_python` compares operation-specific timeout messages. The elapsed duration in a POST error is measured independently, so the test checks its structure and lower bound rather than requiring two independent requests to have identical elapsed time. Phase-specific connect/read/write/pool deadlines and cached Python transport configuration still need broader parity checks + +Native writes retain two documented correctness safeguards. Python's `async_write_secret` does not invalidate the read cache, so a later read can return a value from before a successful write. `test_public_native_vault_write_invalidates_stale_cached_values` verifies that the native public API returns the updated provider value. Python also overwrites the secret when both the data key and description field are named `description`. `test_public_native_vault_write_rejects_description_overwriting_the_secret` rejects that collision before any request. These corrections do not change Python + +The raw response path does not yet establish complete Vault API parity. SDK authentication payload parsing, argument conversion, connection retry behavior, non-HTTP transport errors, nonstandard JSON encodings during rotation and configuration changes during rotation remain under audit. The catalog stays Python-only + + +The Vault boundary update passed 317 provider and bridge tests, including 191 bridge cases. With the extension unavailable, 35 passed and 156 native-only cases skipped. Seven targeted mutations compiled and failed their regression tests: stripped response envelopes, ignored HTTP 400 failures, skipped replacement equality, skipped current-secret checks, stale write caches, fatal old-secret deletion failures and ignored write-error responses. The restored extension passed again. Five additional differential cases reproduced lossy large-integer error conversion before the raw-JSON correction and pass afterward. Live public native reads, writes, deletes and same-name/new-alias rotations passed against local Vault 1.20 with token and AppRole authentication, with Python HTTP construction forbidden diff --git a/litellm-rust/crates/secrets/README.md b/litellm-rust/crates/secrets/README.md index 0d333c7d116..c8b01fe9b3a 100644 --- a/litellm-rust/crates/secrets/README.md +++ b/litellm-rust/crates/secrets/README.md @@ -2,12 +2,42 @@ Construct `SecretManagerState::new(backend, settings)` for a configured manager or use `SecretManagerState::default()` for environment lookups. The configured backend determines its provider identity. Write-only settings and names excluded by `hosted_keys` use the environment directly. `secret_manager_would_be_consulted` follows the same routing decision as resolution -`get_secret` returns `Ok(Some(value))` for a found value, `Ok(None)` when no source contains the value, and `Err(error)` when lookup fails. For managed names, resolution checks the manager, then the environment, then the caller's default. An empty string, `false`, or an explicitly stored JSON null is a found value +Native resolution distinguishes a found value, confirmed absence, and a failed read. Missing values use the caller's default, while provider errors propagate. Empty strings are found values -Backend failures propagate by default. To allow fallback during a backend failure, construct the resolver with `.with_failure_policy(FailurePolicy::EnvironmentFallback)`. It then tries the environment and default, in that order. If neither exists, the original error is returned. This policy applies to manager lookups. Explicit OIDC references retain their own authentication errors and never fall back to environment secrets under the reference name +`new_python_compatible` uses Python's environment fallback and conversion rules. Manager exceptions fall back to the environment, including custom-manager exceptions. AWS missing secrets, failed HTTP reads, missing string payloads, and missing or empty primary secrets return `None` without fallback, matching Python. The standard Python HTTP handler wraps network timeouts in `litellm.Timeout`, which AWS reads also swallow. Invalid primary JSON still raises. An absent AWS primary JSON field and a successful Azure response without a value also remain `None`. Defaults do not replace these results. `.with_failure_policy(FailurePolicy::Propagate)` exposes manager failures explicitly instead. Cancellation and other Python `BaseException`s always propagate unchanged. Explicit OIDC references keep their own errors and never use these fallbacks -`get_secret` preserves value types. `get_secret_str` accepts a string default and rejects boolean or JSON values with `Error::TypeMismatch`. `get_secret_bool` accepts a boolean default and converts strings containing `true` or `false`, ignoring surrounding whitespace and ASCII case. Other strings and JSON values produce `Error::TypeMismatch`. Conversion failures never activate fallback or replace a found value with the default +The getters follow Python's conversion policy. Environment values use case-insensitive, whitespace-trimmed boolean parsing. Manager strings become booleans only when Python literal evaluation yields a boolean; other strings retain their exact contents. Non-string manager results produce `None`. `get_secret_str` returns only strings, and `get_secret_bool` accepts booleans or strings containing `true` or `false`. A type mismatch returns `None` and does not activate fallback -Provider payloads remain strings unless explicitly selecting a field from an AWS primary JSON secret. Google caches only successfully decoded string payloads, so reads have identical values and types before and after caching. Confirmed absence and failed reads are not cached. AWS resource-not-found responses and Google HTTP 404 responses indicate absence. Other provider errors remain errors, and successful responses without the required payload are malformed responses rather than missing secrets +## Route integration + +Inject `Arc` from `litellm_secrets::source` into route preparation. `SecretResolver` implements this interface and supports arbitrary names through its asynchronous `get_secret_str`. It applies the same manager selection, conversion, and fallback policy to every lookup + +For synchronous provider transformations, call `source.resolve(names).await` during preparation and inject the returned `Secrets` snapshot. A snapshot contains only those names and never reads the process environment implicitly. Resolve runtime names through the source before invoking a synchronous transformation. OCR uses this pattern; other routes can adopt it as they are implemented + +The Python bridge uses this shared source and resolver. The shared proxy initializer captures effective configuration, and directly constructed LiteLLM managers are adapted at the dispatch boundary, and the bridge retains a native backend per configured client. Python reads and Rust routes share that backend. Custom Python implementations remain external callbacks. Rollout policy controls whether the native binding is selected. Public provider reads and Vault/CyberArk mutations use this selection; AWS mutation bindings remain unfinished. Mutations update the same native cache used by public reads. Vault retains complete write response bodies and Python-compatible error dictionaries while verifying rotation with fresh reads. The Python boundary retains primary JSON until return conversion so Python JSON numbers, values, and exception details survive unchanged + +## Backend contracts + +AWS Secrets Manager, Azure Key Vault, Google Secret Manager, Vault, and CyberArk implement `BaseSecretManager` for reads with an operation context. Foreign provider contexts are rejected before cache access or I/O. Writes and deletes use separate `SecretWriter` and `SecretDeleter` capabilities. CyberArk rotation writes and verifies the replacement while retaining the old alias because Conjur does not support deletion through this API + +Shared rotation verifies that the replacement has the requested value before deleting the old secret. Same-name rotation keeps the replacement. AWS same-name rotation uses its version update API directly, matching Python + +Backend reads preserve payload strings. Conversion belongs to the resolver. Google caches only successfully decoded payloads, so values agree before and after caching. Native reads do not cache absence or failures. Python-compatible Google reads preserve Python's negative cache and its always-read override. Resource-not-found responses indicate absence; authentication, permission, transport, and malformed successful responses remain errors The HashiCorp Vault backend is enabled with the `hashicorp` feature and reads KV v2 values from `HCP_VAULT_*` environment variables. It supports static tokens, AppRole authentication, and TLS certificate authentication + +## Intentional differences from Python + +Native backends consistently distinguish absence from failure instead of swallowing provider errors. Python-compatible resolution maps these results back to the Python handler contract before applying fallback + +`hosted_keys` excludes a name for every backend. Python's handler recognizes Azure `SecretClient` and Google `KeyManagementServiceClient` instances before the `local` branch, allowing excluded names to reach those providers. Rust treats that as a routing bug. `test_rust_hosted_keys_exclude_azure_sdk_clients_too` in `tests/test_litellm/rust_bridge/ocr/test_secrets.py` pins this behavior + +Google rejects malformed base64 and mismatched CRC32C values instead of accepting corrupted payloads. Python currently ignores the checksum and uses permissive base64 decoding. Rust follows [RFC 4648](https://www.rfc-editor.org/rfc/rfc4648#section-3.3) and [Google's integrity guidance](https://docs.cloud.google.com/secret-manager/docs/data-integrity); `failed_or_missing_reads_are_not_cached` covers rejection and recovery + +CyberArk cached reads preserve the original secret text. Python's shared cache attempts JSON decoding, so a secret such as `"password"` changes to `password` after the first read, and `true` changes to a Boolean. This corrupts the stored credential representation. `test_public_cyberark_reads_reuse_authentication_and_cached_values` demonstrates the Python defect and verifies stable native results + +## Test parity + +[The Python test inventory](PARITY.md) maps each secret-manager test to Rust coverage or its owning boundary + +AWS, Vault, and CyberArk keep client/authentication, reads, and writes/rotation in private provider modules. Their existing integration-test targets group configuration, read/cache, and write/rotation cases, with shared fixtures local to each target. Python-compatible dispatch and string coercion live separately from native dispatch diff --git a/litellm-rust/crates/secrets/src/compatibility.rs b/litellm-rust/crates/secrets/src/compatibility.rs new file mode 100644 index 00000000000..621f7a3f432 --- /dev/null +++ b/litellm-rust/crates/secrets/src/compatibility.rs @@ -0,0 +1,82 @@ +use litellm_core_utils::settings::Lookup; +use litellm_python_compat::{Value, literal::literal_eval}; +use litellm_secrets_types::PythonSecretRead; + +use crate::{ + Error, KeyManagementSettings, KeyManagementSystem, Secret, SecretManager, SecretValue, + get_secret_from_manager, +}; + +pub async fn get_secret_from_python_manager( + manager: &SecretManager, + name: &str, + settings: &KeyManagementSettings, + environment: &(dyn Lookup + Send + Sync), +) -> Result, Error> { + #[cfg(feature = "aws")] + if let SecretManager::AwsSecretsManagerV2(client) = manager { + return client + .read_secret_for_python(name, settings.primary_secret_name.as_deref(), environment) + .await + .map_err(Error::from); + } + let result = match manager { + #[cfg(feature = "cyberark")] + SecretManager::Cyberark(client) => Ok(client + .read_with_retry( + name, + &Default::default(), + litellm_secrets_cyberark::AuthenticationRetry::Never, + ) + .await + .unwrap_or(None) + .map(Secret::String)), + #[cfg(feature = "google")] + SecretManager::GoogleSecretManager(client) => client + .get_secret_for_python(name) + .await + .map_err(Error::from), + _ => get_secret_from_manager(manager, name, settings, environment).await, + }; + match result { + #[cfg(feature = "google")] + Err(Error::Google(litellm_secrets_google::Error::Status(404))) => { + Err(Error::ManagedSecretMissing) + } + #[cfg(feature = "azure")] + Err(Error::Azure(litellm_secrets_azure::Error::MissingValue)) => Ok(None), + Ok(None) + if matches!(manager, SecretManager::External(_)) + && manager.system() == KeyManagementSystem::AzureKeyVault => + { + Ok(None) + } + Ok(None) => Err(Error::ManagedSecretMissing), + result => result, + } +} + +pub(crate) fn python_manager_string(value: SecretValue) -> Secret { + match literal_eval(value.expose()) { + Ok(Value::Bool(boolean)) => Secret::Bool(boolean), + _ => Secret::String(value), + } +} + +pub async fn read_secret_from_python_manager( + manager: &SecretManager, + name: &str, + settings: &KeyManagementSettings, + environment: &(dyn Lookup + Send + Sync), +) -> Result { + #[cfg(feature = "aws")] + if let SecretManager::AwsSecretsManagerV2(client) = manager { + return client + .read_payload_for_python(name, settings.primary_secret_name.as_deref(), environment) + .await + .map_err(Error::from); + } + get_secret_from_python_manager(manager, name, settings, environment) + .await + .map(PythonSecretRead::Value) +} diff --git a/litellm-rust/crates/secrets/src/error.rs b/litellm-rust/crates/secrets/src/error.rs index de325ff4981..07f2f205bec 100644 --- a/litellm-rust/crates/secrets/src/error.rs +++ b/litellm-rust/crates/secrets/src/error.rs @@ -1,5 +1,9 @@ #[derive(Debug, thiserror::Error)] pub enum Error { + #[error("configured secret manager did not return a secret")] + ManagedSecretMissing, + #[error("native secret backend is unavailable for this system")] + NativeBackendUnavailable, #[error("encrypted environment value is missing")] MissingCiphertext, #[error("ciphertext is not valid base64 for the configured manager")] @@ -26,6 +30,8 @@ pub enum Error { TypeMismatch { expected: &'static str }, #[error("external secret manager failed")] ExternalManager(#[source] Box), + #[error("external secret manager read failed")] + ExternalRead(#[source] Box), #[cfg(feature = "aws")] #[error(transparent)] Aws(#[from] litellm_secrets_aws::Error), diff --git a/litellm-rust/crates/secrets/src/handler.rs b/litellm-rust/crates/secrets/src/handler.rs index 0ae84821883..db37d793384 100644 --- a/litellm-rust/crates/secrets/src/handler.rs +++ b/litellm-rust/crates/secrets/src/handler.rs @@ -7,6 +7,14 @@ use crate::{Error, KeyManagementSettings, KeyManagementSystem, Secret}; #[cfg(any(feature = "aws", feature = "google"))] use crate::SecretValue; +#[cfg(any( + feature = "google", + feature = "hashicorp", + feature = "azure", + feature = "cyberark" +))] +use litellm_secrets_types::BaseSecretManager; + pub trait ExternalSecretManager: Send + Sync { fn system(&self) -> KeyManagementSystem; @@ -103,26 +111,13 @@ pub async fn get_secret_from_manager( .await .map_err(Error::from), #[cfg(feature = "google")] - SecretManager::GoogleSecretManager(client) => client - .get_secret_from_google_secret_manager(secret_name) - .await - .map_err(Error::from), + SecretManager::GoogleSecretManager(client) => read_manager(client, secret_name).await, #[cfg(feature = "hashicorp")] - SecretManager::HashicorpVault(client) => client - .async_read_secret(secret_name) - .await - .map(|value| value.map(Secret::String)) - .map_err(Error::from), + SecretManager::HashicorpVault(client) => read_manager(client, secret_name).await, #[cfg(feature = "azure")] - SecretManager::AzureKeyVault(client) => { - client.get_secret(secret_name).await.map_err(Error::from) - } + SecretManager::AzureKeyVault(client) => read_manager(client, secret_name).await, #[cfg(feature = "cyberark")] - SecretManager::Cyberark(client) => client - .async_read_secret(secret_name) - .await - .map(|value| value.map(Secret::String)) - .map_err(Error::from), + SecretManager::Cyberark(client) => read_manager(client, secret_name).await, } } @@ -160,3 +155,23 @@ fn decode_ciphertext(value: &str, mode: Base64Mode) -> Result, Error> { } Ok(ciphertext) } + +#[cfg(any( + feature = "google", + feature = "hashicorp", + feature = "azure", + feature = "cyberark" +))] +async fn read_manager( + manager: &M, + name: &str, +) -> Result, Error> +where + Error: From, +{ + manager + .async_read_secret(name, &M::Context::default()) + .await + .map(|value| value.map(Secret::String)) + .map_err(Error::from) +} diff --git a/litellm-rust/crates/secrets/src/lib.rs b/litellm-rust/crates/secrets/src/lib.rs index 58aba8494fd..8d0f0561608 100644 --- a/litellm-rust/crates/secrets/src/lib.rs +++ b/litellm-rust/crates/secrets/src/lib.rs @@ -1,18 +1,23 @@ #![forbid(unsafe_code)] +mod compatibility; mod error; mod handler; +mod native; mod oidc; mod resolver; +pub mod source; mod state; +pub use compatibility::{get_secret_from_python_manager, read_secret_from_python_manager}; pub use error::Error; pub use handler::{ExternalSecretManager, SecretManager, get_secret_from_manager}; pub use litellm_secrets_types::{ AccessMode, KeyManagementSettings, KeyManagementSystem, Secret, SecretValue, }; +pub use native::load_native_manager; pub use oidc::{OidcProvider, OidcReference, OidcResolver}; -pub use resolver::{FailurePolicy, SecretResolver}; +pub use resolver::{FailurePolicy, SecretResolver, normalize_nonempty_secret_str}; pub use state::{SecretManagerState, secret_manager_would_be_consulted}; #[cfg(feature = "aws")] diff --git a/litellm-rust/crates/secrets/src/native.rs b/litellm-rust/crates/secrets/src/native.rs new file mode 100644 index 00000000000..80f0e46245c --- /dev/null +++ b/litellm-rust/crates/secrets/src/native.rs @@ -0,0 +1,61 @@ +use std::sync::Arc; + +use litellm_core_utils::settings::Lookup; + +use crate::{Error, KeyManagementSettings, KeyManagementSystem, SecretManager}; + +pub async fn load_native_manager( + system: KeyManagementSystem, + settings: KeyManagementSettings, + environment: Arc, + enterprise_enabled: bool, +) -> Result { + match (system, settings, environment, enterprise_enabled) { + #[cfg(feature = "aws")] + (KeyManagementSystem::AwsSecretManager, settings, environment, _) => { + crate::aws::AwsSecretsManagerV2::load_aws_secret_manager( + Some(true), + settings, + environment, + )? + .map(SecretManager::AwsSecretsManagerV2) + .ok_or(Error::NativeBackendUnavailable) + } + #[cfg(feature = "aws")] + (KeyManagementSystem::AwsKms, settings, environment, _) => { + crate::aws::load_aws_kms(Some(true), &settings, environment)? + .map(SecretManager::AwsKms) + .ok_or(Error::NativeBackendUnavailable) + } + #[cfg(feature = "azure")] + (KeyManagementSystem::AzureKeyVault, _, environment, _) => Ok( + SecretManager::AzureKeyVault(crate::azure::AzureKeyVault::new(environment)?), + ), + #[cfg(feature = "google")] + (KeyManagementSystem::GoogleSecretManager, _, environment, enterprise_enabled) => { + Ok(SecretManager::GoogleSecretManager( + crate::google::GoogleSecretManager::new(environment, enterprise_enabled)?, + )) + } + #[cfg(feature = "google")] + (KeyManagementSystem::GoogleKms, _, environment, _) => { + crate::google::load_google_kms(Some(true), environment) + .await? + .map(SecretManager::GoogleKms) + .ok_or(Error::NativeBackendUnavailable) + } + #[cfg(feature = "hashicorp")] + (KeyManagementSystem::HashicorpVault, _, environment, enterprise_enabled) => { + Ok(SecretManager::HashicorpVault( + crate::hashicorp::HashicorpVault::new(environment, enterprise_enabled)?, + )) + } + #[cfg(feature = "cyberark")] + (KeyManagementSystem::Cyberark, _, environment, enterprise_enabled) => { + Ok(SecretManager::Cyberark( + crate::cyberark::CyberArkSecretManager::new(environment, enterprise_enabled)?, + )) + } + _ => Err(Error::NativeBackendUnavailable), + } +} diff --git a/litellm-rust/crates/secrets/src/oidc.rs b/litellm-rust/crates/secrets/src/oidc.rs index fd477859bf6..f3c1e38ce7b 100644 --- a/litellm-rust/crates/secrets/src/oidc.rs +++ b/litellm-rust/crates/secrets/src/oidc.rs @@ -3,7 +3,7 @@ use std::{ time::{Duration, SystemTime, UNIX_EPOCH}, }; -use jsonwebtoken::dangerous::insecure_decode_claims; +use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD}; use litellm_core_utils::settings::Lookup; use moka::future::Cache; use serde::Deserialize; @@ -67,13 +67,15 @@ struct OidcTokenClaims { enum NumericDate { Number(f64), String(String), + Boolean(bool), } impl NumericDate { fn seconds(self) -> Option { match self { Self::Number(value) => Some(value), - Self::String(value) => value.parse().ok(), + Self::String(value) => value.trim().parse().ok(), + Self::Boolean(value) => Some(f64::from(u8::from(value))), } .filter(|value| value.is_finite()) } @@ -84,6 +86,8 @@ pub struct OidcResolver { google_identity_endpoint: reqwest::Url, cache: Cache, clock: fn() -> SystemTime, + #[cfg(feature = "azure")] + azure_token_provider: std::sync::Arc, } impl Default for OidcResolver { @@ -110,6 +114,21 @@ impl OidcResolver { .time_to_live(GOOGLE_TOKEN_MAX_TTL) .build(), clock: SystemTime::now, + #[cfg(feature = "azure")] + azure_token_provider: std::sync::Arc::new( + litellm_secrets_azure::NativeAzureTokenProvider::default(), + ), + } + } + + #[cfg(feature = "azure")] + pub fn with_azure_token_provider( + self, + provider: std::sync::Arc, + ) -> Self { + Self { + azure_token_provider: provider, + ..self } } @@ -141,6 +160,15 @@ impl OidcResolver { if let Some(path) = environment.get(AZURE_FEDERATED_TOKEN_FILE) { return read_file(&path).await.map(Some); } + #[cfg(feature = "azure")] + { + self.azure_token_provider + .get_token(audience, environment) + .await + .map(Some) + .map_err(Error::Azure) + } + #[cfg(not(feature = "azure"))] Err(Error::UnsupportedOidc) } OidcProvider::Github => { @@ -256,7 +284,14 @@ async fn read_allowed_file( fn oidc_token_cache_ttl(token: &str, now: SystemTime, max_ttl: Duration) -> Option { let fallback = Some(max_ttl); - let Ok(claims) = insecure_decode_claims::(token) else { + let segments: Vec<_> = token.split('.').collect(); + let [_, payload, _] = segments.as_slice() else { + return fallback; + }; + let Ok(decoded) = URL_SAFE_NO_PAD.decode(payload.trim_end_matches('=')) else { + return fallback; + }; + let Ok(claims) = serde_json::from_slice::(&decoded) else { return fallback; }; let Some(exp) = claims.exp.and_then(NumericDate::seconds) else { diff --git a/litellm-rust/crates/secrets/src/resolver.rs b/litellm-rust/crates/secrets/src/resolver.rs index 597ca11b171..69445e5b410 100644 --- a/litellm-rust/crates/secrets/src/resolver.rs +++ b/litellm-rust/crates/secrets/src/resolver.rs @@ -1,6 +1,10 @@ use std::sync::Arc; -use litellm_core_utils::settings::{Lookup, ProcessEnvironment}; +use crate::compatibility::python_manager_string; +use litellm_core_utils::{ + serde_compat::parse_str_bool, + settings::{Lookup, ProcessEnvironment}, +}; use crate::state::{LookupTarget, normalize_secret_name}; use crate::{Error, OidcResolver, Secret, SecretManagerState, SecretValue}; @@ -17,6 +21,7 @@ pub struct SecretResolver { environment: Arc, oidc: OidcResolver, failure_policy: FailurePolicy, + python_compatible: bool, } impl Default for SecretResolver { @@ -40,6 +45,19 @@ impl SecretResolver { environment, oidc, failure_policy: FailurePolicy::default(), + python_compatible: false, + } + } + + pub fn new_python_compatible( + state: Arc, + environment: Arc, + oidc: OidcResolver, + ) -> Self { + Self { + python_compatible: true, + failure_policy: FailurePolicy::EnvironmentFallback, + ..Self::new(state, environment, oidc) } } @@ -54,6 +72,19 @@ impl SecretResolver { &self, name: &str, default_value: Option, + ) -> Result, Error> { + let value = self.read(name, default_value.clone()).await?; + Ok(if self.python_compatible { + value + } else { + value.or(default_value) + }) + } + + async fn read( + &self, + name: &str, + default_value: Option, ) -> Result, Error> { let name = normalize_secret_name(name); if name.starts_with("oidc/") { @@ -61,36 +92,38 @@ impl SecretResolver { .oidc .resolve(name, self.environment.as_ref()) .await - .map(|value| value.map(Secret::String).or(default_value)); + .map(|value| value.map(Secret::String)); } let LookupTarget::Manager { backend, settings } = self.state.lookup_target(name) else { - return Ok(self.environment_secret(name).or(default_value)); + return Ok(self.environment_value(name)); }; - match crate::get_secret_from_manager(backend, name, settings, self.environment.as_ref()) + let result = if self.python_compatible { + crate::get_secret_from_python_manager( + backend, + name, + settings, + self.environment.as_ref(), + ) .await - { - Ok(value) => Ok(value - .or_else(|| self.environment_secret(name)) - .or(default_value)), + } else { + crate::get_secret_from_manager(backend, name, settings, self.environment.as_ref()).await + }; + match result { + Ok(value) => Ok(value.and_then(|value| self.manager_value(value))), Err(error @ Error::ExternalManager(_)) => Err(error), Err(error) => match self.failure_policy { + FailurePolicy::Propagate if self.python_compatible => { + default_value.map(Some).ok_or(error) + } FailurePolicy::Propagate => Err(error), - FailurePolicy::EnvironmentFallback => self - .environment_secret(name) - .or(default_value) - .map(Some) - .ok_or(error), + FailurePolicy::EnvironmentFallback => Ok(self + .environment + .get(name) + .and_then(|value| self.manager_value(Secret::String(SecretValue::new(value))))), }, } } - fn environment_secret(&self, name: &str) -> Option { - self.environment - .get(name) - .map(SecretValue::new) - .map(Secret::String) - } - pub async fn get_secret_str( &self, name: &str, @@ -102,6 +135,7 @@ impl SecretResolver { { Some(Secret::String(value)) => Ok(Some(value)), None => Ok(None), + Some(Secret::Bool(_) | Secret::Json(_)) if self.python_compatible => Ok(None), Some(Secret::Bool(_) | Secret::Json(_)) => { Err(Error::TypeMismatch { expected: "string" }) } @@ -118,19 +152,57 @@ impl SecretResolver { .await? { Some(Secret::Bool(value)) => Ok(Some(value)), - Some(Secret::String(value)) => { - match value.expose().trim().to_ascii_lowercase().as_str() { - "true" => Ok(Some(true)), - "false" => Ok(Some(false)), - _ => Err(Error::TypeMismatch { - expected: "boolean", - }), - } - } + Some(Secret::String(value)) => match parse_str_bool(value.expose()) { + Some(value) => Ok(Some(value)), + None if self.python_compatible => Ok(None), + None => Err(Error::TypeMismatch { + expected: "boolean", + }), + }, + Some(Secret::Json(_)) if self.python_compatible => Ok(None), Some(Secret::Json(_)) => Err(Error::TypeMismatch { expected: "boolean", }), None => Ok(None), } } + + fn environment_value(&self, name: &str) -> Option { + let value = self.environment.get(name)?; + if !self.python_compatible { + return Some(Secret::String(SecretValue::new(value))); + } + if self + .state + .settings() + .is_some_and(|settings| settings.access_mode.readable()) + { + return Some(python_manager_string(SecretValue::new(value))); + } + Some( + parse_str_bool(&value) + .map_or_else(|| Secret::String(SecretValue::new(value)), Secret::Bool), + ) + } + + fn manager_value(&self, secret: Secret) -> Option { + if !self.python_compatible { + return Some(secret); + } + + let Secret::String(value) = secret else { + return None; + }; + Some(python_manager_string(value)) + } +} + +pub fn normalize_nonempty_secret_str(value: Option<&str>) -> Option<&str> { + value + .map(|value| { + value.trim_matches(|character: char| { + character.is_whitespace() || matches!(character, '\u{1c}'..='\u{1f}') + }) + }) + .filter(|value| !value.is_empty()) } diff --git a/litellm-rust/crates/secrets/src/source.rs b/litellm-rust/crates/secrets/src/source.rs new file mode 100644 index 00000000000..a1615bad055 --- /dev/null +++ b/litellm-rust/crates/secrets/src/source.rs @@ -0,0 +1,73 @@ +use std::{collections::HashMap, sync::Arc}; + +use futures_util::future::{BoxFuture, try_join_all}; +use litellm_core_utils::settings::Lookup; + +use crate::{Error, SecretResolver, SecretValue}; + +pub type Secrets = Arc; + +pub trait SecretSource: Send + Sync { + fn get_secret_str<'a>( + &'a self, + name: &'a str, + ) -> BoxFuture<'a, Result, Error>>; + + fn resolve<'a>(&'a self, names: &'a [&str]) -> BoxFuture<'a, Result> { + Box::pin(async move { + let values = try_join_all(names.iter().map(|name| async move { + self.get_secret_str(name) + .await + .map(|value| ((*name).to_owned(), value)) + })) + .await? + .into_iter() + .collect(); + Ok(Arc::new(SecretSnapshot { values }) as Secrets) + }) + } +} + +impl SecretSource for SecretResolver { + fn get_secret_str<'a>( + &'a self, + name: &'a str, + ) -> BoxFuture<'a, Result, Error>> { + Box::pin(SecretResolver::get_secret_str(self, name, None)) + } +} + +#[derive(Default)] +pub struct EnvironmentSecrets(SecretResolver); + +impl EnvironmentSecrets { + pub fn python_compatible() -> Self { + Self(SecretResolver::new_python_compatible( + Arc::new(crate::SecretManagerState::default()), + Arc::new(litellm_core_utils::settings::ProcessEnvironment), + crate::OidcResolver::default(), + )) + } +} + +impl SecretSource for EnvironmentSecrets { + fn get_secret_str<'a>( + &'a self, + name: &'a str, + ) -> BoxFuture<'a, Result, Error>> { + Box::pin(self.0.get_secret_str(name, None)) + } +} + +struct SecretSnapshot { + values: HashMap>, +} + +impl Lookup for SecretSnapshot { + fn get(&self, name: &str) -> Option { + self.values + .get(name) + .and_then(Option::as_ref) + .map(|value| value.expose().to_owned()) + } +} diff --git a/litellm-rust/crates/secrets/tests/aws.rs b/litellm-rust/crates/secrets/tests/aws.rs new file mode 100644 index 00000000000..174d0881339 --- /dev/null +++ b/litellm-rust/crates/secrets/tests/aws.rs @@ -0,0 +1,223 @@ +#![cfg(feature = "aws")] + +use std::sync::Arc; + +use litellm_secrets::{ + AccessMode, Error, FailurePolicy, KeyManagementSettings, OidcResolver, Secret, SecretManager, + SecretManagerState, SecretResolver, SecretValue, aws::AwsSecretsManagerV2, + secret_manager_would_be_consulted, +}; +use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method}; + +fn state(server: &MockServer, settings: KeyManagementSettings) -> SecretManagerState { + let endpoint = server.uri(); + let environment = Arc::new(move |name: &str| match name { + "AWS_REGION_NAME" => Some("us-east-1".into()), + "AWS_ACCESS_KEY_ID" | "AWS_SECRET_ACCESS_KEY" => Some("test".into()), + "AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(endpoint.clone()), + _ => None, + }); + let manager = + AwsSecretsManagerV2::load_aws_secret_manager(Some(true), settings.clone(), environment) + .unwrap() + .unwrap(); + SecretManagerState::new(SecretManager::AwsSecretsManagerV2(manager), settings) +} + +#[rstest::rstest] +#[case::missing(400, serde_json::json!({"__type":"ResourceNotFoundException"}), None)] +#[case::denied(400, serde_json::json!({"__type":"AccessDeniedException"}), None)] +#[case::malformed(200, serde_json::json!({}), None)] +#[case::invalid_primary(200, serde_json::json!({"SecretString":"not-json"}), Some("primary"))] +#[tokio::test] +async fn read_results_follow_the_selected_failure_policy( + #[case] status: u16, + #[case] body: serde_json::Value, + #[case] primary_secret_name: Option<&str>, + #[values(FailurePolicy::Propagate, FailurePolicy::EnvironmentFallback)] policy: FailurePolicy, + #[values(None, Some("environment"))] environment: Option<&'static str>, + #[values(None, Some("default"))] default: Option<&str>, +) { + let server = MockServer::start().await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(status).set_body_json(body.clone())) + .expect(1) + .mount(&server) + .await; + let resolver = SecretResolver::new_python_compatible( + Arc::new(state( + &server, + KeyManagementSettings { + primary_secret_name: primary_secret_name.map(str::to_owned), + ..Default::default() + }, + )), + Arc::new(move |_: &str| environment.map(str::to_owned)), + OidcResolver::default(), + ) + .with_failure_policy(policy); + let result = resolver + .get_secret_str("KEY", default.map(SecretValue::new)) + .await; + if primary_secret_name.is_none() { + assert_eq!(result.unwrap(), None); + } else if policy == FailurePolicy::EnvironmentFallback { + assert_eq!( + result.unwrap().as_ref().map(SecretValue::expose), + environment + ); + } else if let Some(default) = default { + assert_eq!(result.unwrap().unwrap().expose(), default); + } else { + assert!(matches!(result, Err(Error::Aws(_)))); + } +} + +#[rstest::rstest] +#[case::boolean(serde_json::json!(false))] +#[case::object(serde_json::json!({"key":1}))] +#[case::null(serde_json::Value::Null)] +#[case::string(serde_json::json!("true"))] +#[tokio::test] +async fn primary_secret_values_other_than_strings_resolve_to_none( + #[case] value: serde_json::Value, +) { + let server = MockServer::start().await; + Mock::given(method("POST")) + .respond_with(ResponseTemplate::new(200).set_body_json( + serde_json::json!({"SecretString":serde_json::json!({"KEY":value}).to_string()}), + )) + .expect(3) + .mount(&server) + .await; + let settings = KeyManagementSettings { + primary_secret_name: Some("primary".into()), + ..Default::default() + }; + let resolver = SecretResolver::new_python_compatible( + Arc::new(state(&server, settings)), + Arc::new(|_: &str| Some("fallback".into())), + OidcResolver::default(), + ); + let text = value.as_str(); + assert_eq!( + resolver + .get_secret("KEY", Some(Secret::Bool(true))) + .await + .unwrap(), + text.map(|text| Secret::String(SecretValue::new(text))) + ); + assert_eq!( + resolver + .get_secret_str("KEY", None) + .await + .unwrap() + .as_ref() + .map(SecretValue::expose), + text + ); + assert_eq!( + resolver.get_secret_bool("KEY", None).await.unwrap(), + text.map(|_| true) + ); +} + +#[rstest::rstest] +#[tokio::test] +async fn gating_prediction_matches_actual_lookup( + #[values(AccessMode::ReadOnly, AccessMode::WriteOnly, AccessMode::ReadAndWrite)] + access_mode: AccessMode, + #[values(None, Some(vec![]), Some(vec!["KEY".into()]))] hosted_keys: Option>, + #[values("os.environ/KEY", "os.environ/oidc/env/KEY")] name: &str, +) { + let server = MockServer::start().await; + let expected = name == "os.environ/KEY" + && access_mode.readable() + && hosted_keys + .as_ref() + .is_none_or(|keys| keys.iter().any(|key| key == "KEY")); + Mock::given(method("POST")) + .respond_with( + ResponseTemplate::new(200).set_body_json(serde_json::json!({"SecretString":"remote"})), + ) + .expect(u64::from(expected)) + .mount(&server) + .await; + let state = state( + &server, + KeyManagementSettings { + access_mode, + hosted_keys, + ..Default::default() + }, + ); + assert!(state.backend().is_some()); + assert_eq!(state.settings().unwrap().access_mode, access_mode); + assert_eq!(secret_manager_would_be_consulted(&state, name), expected); + let resolver = SecretResolver::new_python_compatible( + Arc::new(state), + Arc::new(|_: &str| Some("environment".into())), + OidcResolver::default(), + ); + assert_eq!( + resolver + .get_secret_str(name, None) + .await + .unwrap() + .unwrap() + .expose(), + if expected { "remote" } else { "environment" } + ); +} + +#[tokio::test] +async fn aws_handler_reads_ciphertext_decodes_trims_and_redacts() { + use aws_sdk_kms::{ + Client, + config::{BehaviorVersion, Credentials, Region}, + }; + use base64::{Engine, engine::general_purpose::STANDARD}; + use litellm_secrets::{ + Error, KeyManagementSettings, SecretManager, aws::AwsKms, get_secret_from_manager, + }; + use wiremock::{Mock, MockServer, ResponseTemplate, matchers::body_json}; + + let server = MockServer::start().await; + Mock::given(body_json( + serde_json::json!({"CiphertextBlob": STANDARD.encode("encrypted")}), + )) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(serde_json::json!({"Plaintext":STANDARD.encode(" value\n")})), + ) + .expect(1) + .mount(&server) + .await; + let client = Client::from_conf( + aws_sdk_kms::Config::builder() + .behavior_version(BehaviorVersion::latest()) + .region(Region::new("us-east-1")) + .credentials_provider(Credentials::new("test", "test", None, None, "test")) + .endpoint_url(server.uri()) + .build(), + ); + let manager = SecretManager::AwsKms(AwsKms::new(client)); + let settings = KeyManagementSettings::default(); + let value = get_secret_from_manager(&manager, "KEY", &settings, &|name: &str| { + assert_eq!(name, "KEY"); + Some(format!(" {}\n", STANDARD.encode("encrypted"))) + }) + .await + .unwrap() + .unwrap(); + assert_eq!(value.as_str(), Some("value")); + assert!(!format!("{value:?}").contains("value")); + assert!(matches!( + get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| None).await, + Err(Error::MissingCiphertext) + )); + assert!(matches!( + get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| Some("abc".into())).await, + Err(Error::InvalidCiphertext) + )); +} diff --git a/litellm-rust/crates/secrets/tests/azure.rs b/litellm-rust/crates/secrets/tests/azure.rs new file mode 100644 index 00000000000..b844b198cd4 --- /dev/null +++ b/litellm-rust/crates/secrets/tests/azure.rs @@ -0,0 +1,106 @@ +#![cfg(feature = "azure")] + +#[tokio::test] +async fn azure_handler_reads_missing_and_failed_secrets() { + use litellm_secrets::{ + Error, KeyManagementSettings, KeyManagementSystem, SecretManager, azure::AzureKeyVault, + get_secret_from_manager, + }; + use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{path, query_param}, + }; + + let server = MockServer::start().await; + Mock::given(path("/secrets/KEY")) + .and(query_param("api-version", "7.4")) + .respond_with( + ResponseTemplate::new(200).set_body_json(serde_json::json!({"value": "value"})), + ) + .expect(1) + .mount(&server) + .await; + let manager = SecretManager::AzureKeyVault( + AzureKeyVault::with_client( + reqwest::Client::new(), + server.uri().parse().unwrap(), + std::sync::Arc::new(|name: &str| (name == "AZURE_AD_TOKEN").then(|| "fake".to_owned())), + ) + .unwrap(), + ); + assert_eq!(manager.system(), KeyManagementSystem::AzureKeyVault); + let settings = KeyManagementSettings::default(); + let value = get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| None) + .await + .unwrap() + .unwrap(); + assert_eq!(value.as_str(), Some("value")); + + let not_found = Mock::given(path("/secrets/MISSING")) + .respond_with(ResponseTemplate::new(404)) + .expect(1) + .mount_as_scoped(&server) + .await; + assert_eq!( + get_secret_from_manager(&manager, "MISSING", &settings, &|_: &str| None) + .await + .unwrap(), + None + ); + drop(not_found); + + Mock::given(path("/secrets/FAILED")) + .respond_with(ResponseTemplate::new(500)) + .expect(1) + .mount(&server) + .await; + assert!(matches!( + get_secret_from_manager(&manager, "FAILED", &settings, &|_: &str| None).await, + Err(Error::Azure(_)) + )); +} + +#[rstest::rstest] +#[case::null(serde_json::json!({"value":null}))] +#[case::absent(serde_json::json!({}))] +#[case::empty(serde_json::json!({"value":""}))] +#[tokio::test] +async fn successful_azure_responses_do_not_fall_back_when_the_value_is_empty_or_null( + #[case] body: serde_json::Value, +) { + use litellm_secrets::{ + OidcResolver, SecretManager, SecretManagerState, SecretResolver, SecretValue, + azure::AzureKeyVault, + }; + use std::sync::Arc; + use wiremock::{Mock, MockServer, ResponseTemplate, matchers::any}; + let server = MockServer::start().await; + Mock::given(any()) + .respond_with(ResponseTemplate::new(200).set_body_json(body.clone())) + .expect(1) + .mount(&server) + .await; + let manager = AzureKeyVault::with_client( + reqwest::Client::new(), + server.uri().parse().unwrap(), + Arc::new(|name: &str| (name == "AZURE_AD_TOKEN").then(|| "token".into())), + ) + .unwrap(); + let resolver = SecretResolver::new_python_compatible( + Arc::new(SecretManagerState::new( + SecretManager::AzureKeyVault(manager), + Default::default(), + )), + Arc::new(|_: &str| Some("environment".into())), + OidcResolver::default(), + ); + assert_eq!( + resolver + .get_secret_str("KEY", Some(SecretValue::new("default"))) + .await + .unwrap() + .as_ref() + .map(SecretValue::expose), + body.get("value").and_then(serde_json::Value::as_str) + ); +} diff --git a/litellm-rust/crates/secrets/tests/common_read_contract.rs b/litellm-rust/crates/secrets/tests/common_read_contract.rs new file mode 100644 index 00000000000..698e1ad8f63 --- /dev/null +++ b/litellm-rust/crates/secrets/tests/common_read_contract.rs @@ -0,0 +1,260 @@ +#![cfg(all( + feature = "aws", + feature = "azure", + feature = "google", + feature = "hashicorp", + feature = "cyberark" +))] +use std::sync::Arc; + +use base64::{Engine, engine::general_purpose::STANDARD}; +use litellm_core_utils::settings::Lookup; +use litellm_secrets::{ + KeyManagementSettings, SecretManager, SecretValue, + aws::AwsSecretsManagerV2, + azure::AzureKeyVault, + cyberark::CyberArkSecretManager, + get_secret_from_manager, + google::GoogleSecretManager, + hashicorp::{HashicorpVault, HashicorpVaultConfig}, +}; +use rstest::rstest; +use serde_json::json; +use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{any, path}, +}; + +#[derive(Clone, Copy, Debug)] +enum Provider { + Aws, + Azure, + Google, + Vault, + Cyberark, +} + +fn manager(provider: Provider, server: &MockServer) -> SecretManager { + let environment: Arc = Arc::new({ + let address = server.uri(); + move |name: &str| match name { + "HCP_VAULT_ADDR" | "AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(address.clone()), + "AWS_ACCESS_KEY_ID" | "AWS_SECRET_ACCESS_KEY" => Some("test".into()), + "AZURE_AD_TOKEN" | "VERTEX_AI_API_KEY" | "HCP_VAULT_TOKEN" => Some("token".into()), + _ => None, + } + }); + match provider { + Provider::Aws => SecretManager::AwsSecretsManagerV2( + AwsSecretsManagerV2::load_aws_secret_manager( + Some(true), + KeyManagementSettings { + aws_region_name: Some("us-east-1".into()), + ..Default::default() + }, + environment, + ) + .unwrap() + .unwrap(), + ), + Provider::Azure => SecretManager::AzureKeyVault( + AzureKeyVault::with_client( + reqwest::Client::new(), + server.uri().parse().unwrap(), + environment, + ) + .unwrap(), + ), + Provider::Google => SecretManager::GoogleSecretManager( + GoogleSecretManager::with_client( + reqwest::Client::new(), + server.uri().parse().unwrap(), + "project".into(), + environment, + None, + false, + ) + .unwrap(), + ), + Provider::Vault => SecretManager::HashicorpVault( + HashicorpVault::from_config( + HashicorpVaultConfig::from_environment(environment.as_ref()).unwrap(), + true, + ) + .unwrap(), + ), + Provider::Cyberark => SecretManager::Cyberark(CyberArkSecretManager::with_client( + reqwest::Client::new(), + server.uri().parse().unwrap(), + "acct".into(), + "admin".into(), + SecretValue::new("key"), + None, + )), + } +} + +fn response(provider: Provider, value: &str) -> ResponseTemplate { + match provider { + Provider::Aws => ResponseTemplate::new(200).set_body_json(json!({"SecretString": value})), + Provider::Azure => ResponseTemplate::new(200).set_body_json(json!({"value": value})), + Provider::Google => ResponseTemplate::new(200) + .set_body_json(json!({"payload": {"data": STANDARD.encode(value)}})), + Provider::Vault => ResponseTemplate::new(200).set_body_json(json!({ + "data": {"data": {"key": value}, "metadata": { + "created_time": "", "deletion_time": "", "custom_metadata": null, + "destroyed": false, "version": 1 + }}, "lease_id": "", "lease_duration": 0, "renewable": false, + "request_id": "", "warnings": null, "wrap_info": null + })), + Provider::Cyberark => ResponseTemplate::new(200).set_body_string(value), + } +} + +#[rstest] +#[case::aws(Provider::Aws)] +#[case::azure(Provider::Azure)] +#[case::google(Provider::Google)] +#[case::vault(Provider::Vault)] +#[case::cyberark(Provider::Cyberark)] +#[tokio::test] +async fn reads_preserve_values_and_distinguish_absence_from_failure(#[case] provider: Provider) { + let server = MockServer::start().await; + Mock::given(path("/authn/acct/admin/authenticate")) + .respond_with(ResponseTemplate::new(200).set_body_string("token")) + .with_priority(1) + .mount(&server) + .await; + let manager = manager(provider, &server); + let settings = KeyManagementSettings::default(); + for (name, value) in [ + ("TEXT", " value\n"), + ("EMPTY", ""), + ("BOOLEAN", "True"), + ("JSON", "{\"key\":1}"), + ] { + let guard = Mock::given(any()) + .respond_with(response(provider, value)) + .with_priority(2) + .mount_as_scoped(&server) + .await; + for _ in 0..2 { + let result = get_secret_from_manager(&manager, name, &settings, &|_: &str| None) + .await + .unwrap() + .unwrap(); + assert_eq!(result.as_str(), Some(value)); + } + drop(guard); + } + let missing = match provider { + Provider::Aws => ResponseTemplate::new(400) + .set_body_json(json!({"__type": "ResourceNotFoundException", "Message": "missing"})), + _ => ResponseTemplate::new(404).set_body_json(json!({"errors": ["missing"]})), + }; + let guard = Mock::given(any()) + .respond_with(missing) + .with_priority(2) + .expect(2) + .mount_as_scoped(&server) + .await; + for _ in 0..2 { + assert!( + get_secret_from_manager(&manager, "MISSING", &settings, &|_: &str| None) + .await + .unwrap() + .is_none() + ); + } + drop(guard); + let guard = Mock::given(any()) + .respond_with(ResponseTemplate::new(403).set_body_json(json!({"errors": ["forbidden"]}))) + .with_priority(2) + .expect(2) + .mount_as_scoped(&server) + .await; + for _ in 0..2 { + assert!( + get_secret_from_manager(&manager, "FAILED", &settings, &|_: &str| None) + .await + .is_err() + ); + } + drop(guard); + let guard = Mock::given(any()) + .respond_with(response(provider, "recovered")) + .with_priority(2) + .expect(2) + .mount_as_scoped(&server) + .await; + for name in ["MISSING", "FAILED"] { + assert_eq!( + get_secret_from_manager(&manager, name, &settings, &|_: &str| None) + .await + .unwrap() + .unwrap() + .as_str(), + Some("recovered") + ); + } + drop(guard); +} + +#[rstest] +#[case::aws(Provider::Aws)] +#[case::azure(Provider::Azure)] +#[case::google(Provider::Google)] +#[case::vault(Provider::Vault)] +#[case::cyberark(Provider::Cyberark)] +#[tokio::test] +async fn python_read_failures_preserve_provider_fallback_rules( + #[case] provider: Provider, + #[values(false, true)] missing: bool, + #[values(None, Some("environment"), Some("True"), Some("true"))] environment_value: Option< + &'static str, + >, +) { + use litellm_secrets::{OidcResolver, Secret, SecretManagerState, SecretResolver}; + let server = MockServer::start().await; + Mock::given(path("/authn/acct/admin/authenticate")) + .respond_with(ResponseTemplate::new(200).set_body_string("token")) + .with_priority(1) + .mount(&server) + .await; + let response = match (provider, missing) { + (Provider::Aws, true) => { + ResponseTemplate::new(400).set_body_json(json!({"__type":"ResourceNotFoundException"})) + } + (_, true) => ResponseTemplate::new(404).set_body_json(json!({"errors":["missing"]})), + (_, false) => ResponseTemplate::new(403).set_body_json(json!({"errors":["forbidden"]})), + }; + Mock::given(any()) + .respond_with(response) + .with_priority(2) + .expect(1) + .mount(&server) + .await; + let resolver = SecretResolver::new_python_compatible( + Arc::new(SecretManagerState::new( + manager(provider, &server), + KeyManagementSettings::default(), + )), + Arc::new(move |_: &str| environment_value.map(str::to_owned)), + OidcResolver::default(), + ); + let expected = if matches!(provider, Provider::Aws) { + None + } else { + environment_value.map(|value| match value { + "True" => Secret::Bool(true), + value => Secret::String(SecretValue::new(value)), + }) + }; + assert_eq!( + resolver + .get_secret("KEY", Some(Secret::String(SecretValue::new("default")))) + .await + .unwrap(), + expected + ); +} diff --git a/litellm-rust/crates/secrets/tests/cyberark.rs b/litellm-rust/crates/secrets/tests/cyberark.rs new file mode 100644 index 00000000000..706c35752d7 --- /dev/null +++ b/litellm-rust/crates/secrets/tests/cyberark.rs @@ -0,0 +1,53 @@ +#![cfg(feature = "cyberark")] + +#[tokio::test] +async fn cyberark_handler_reads_values_and_surfaces_errors() { + use std::time::Duration; + + use litellm_secrets::{ + Error, KeyManagementSettings, SecretManager, SecretValue, cyberark::CyberArkSecretManager, + get_secret_from_manager, + }; + use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{body_string, path}, + }; + + let server = MockServer::start().await; + Mock::given(path("/authn/acct/admin/authenticate")) + .and(body_string("k3y")) + .respond_with(ResponseTemplate::new(200).set_body_string("token")) + .mount(&server) + .await; + Mock::given(path("/secrets/acct/variable/KEY")) + .respond_with(ResponseTemplate::new(200).set_body_string("value")) + .mount(&server) + .await; + let manager = SecretManager::Cyberark(CyberArkSecretManager::with_client( + reqwest::Client::new(), + server.uri().parse().unwrap(), + "acct".into(), + "admin".into(), + SecretValue::new("k3y"), + Some(Duration::from_secs(60)), + )); + assert_eq!( + manager.system(), + litellm_secrets::KeyManagementSystem::Cyberark + ); + let settings = KeyManagementSettings::default(); + let value = get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| None) + .await + .unwrap() + .unwrap(); + assert_eq!(value.as_str(), Some("value")); + + Mock::given(path("/secrets/acct/variable/ERROR")) + .respond_with(ResponseTemplate::new(500)) + .mount(&server) + .await; + assert!(matches!( + get_secret_from_manager(&manager, "ERROR", &settings, &|_: &str| None).await, + Err(Error::Cyberark(_)) + )); +} diff --git a/litellm-rust/crates/secrets/tests/google.rs b/litellm-rust/crates/secrets/tests/google.rs new file mode 100644 index 00000000000..67954fedb78 --- /dev/null +++ b/litellm-rust/crates/secrets/tests/google.rs @@ -0,0 +1,117 @@ +#![cfg(feature = "google")] + +use std::sync::Arc; + +#[rstest::rstest] +#[case::missing(404)] +#[case::failure(503)] +#[tokio::test] +async fn google_resolver_distinguishes_absence_from_failure(#[case] status: u16) { + use litellm_secrets::{ + Error, FailurePolicy, KeyManagementSettings, OidcResolver, SecretManager, + SecretManagerState, SecretResolver, SecretValue, google::GoogleSecretManager, + }; + use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method}; + let server = MockServer::start().await; + Mock::given(method("GET")) + .respond_with(ResponseTemplate::new(status)) + .expect(1) + .mount(&server) + .await; + let environment: Arc = + Arc::new(|name: &str| match name { + "VERTEX_AI_API_KEY" => Some("token".into()), + "KEY" => Some("environment".into()), + _ => None, + }); + let manager = GoogleSecretManager::with_client( + reqwest::Client::new(), + server.uri().parse().unwrap(), + "project".into(), + environment.clone(), + None, + false, + ) + .unwrap(); + let state = SecretManagerState::new( + SecretManager::GoogleSecretManager(manager), + KeyManagementSettings::default(), + ); + let resolver = SecretResolver::new_python_compatible( + Arc::new(state), + environment, + OidcResolver::default(), + ) + .with_failure_policy(FailurePolicy::Propagate); + let result = resolver.get_secret_str("KEY", None).await; + if status == 404 { + assert!(matches!(result, Err(Error::ManagedSecretMissing))); + } else { + assert!( + matches!(result, Err(Error::Google(litellm_secrets::google::Error::Status(actual))) if actual == status) + ); + } + let fallback = resolver + .with_failure_policy(FailurePolicy::EnvironmentFallback) + .get_secret_str("KEY", None) + .await + .unwrap(); + assert_eq!( + fallback.as_ref().map(SecretValue::expose), + Some("environment") + ); +} + +#[tokio::test] +async fn google_handler_requires_canonical_base64_and_preserves_plaintext_whitespace() { + use base64::{Engine, engine::general_purpose::STANDARD}; + use google_cloud_kms_v1::client::KeyManagementService; + use litellm_secrets::{ + Error, KeyManagementSettings, SecretManager, get_secret_from_manager, google::GoogleKms, + }; + use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{body_json, path}, + }; + + let server = MockServer::start().await; + let resource = "projects/project/locations/global/keyRings/ring/cryptoKeys/key"; + Mock::given(path(format!("/v1/{resource}:decrypt"))) + .and(body_json( + serde_json::json!({"ciphertext":STANDARD.encode("encrypted")}), + )) + .respond_with( + ResponseTemplate::new(200) + .set_body_json(serde_json::json!({"plaintext":STANDARD.encode(" value\n")})), + ) + .expect(1) + .mount(&server) + .await; + let client = KeyManagementService::builder() + .with_endpoint(server.uri()) + .with_credentials(google_cloud_auth::credentials::anonymous::Builder::new().build()) + .build() + .await + .unwrap(); + let manager = SecretManager::GoogleKms(GoogleKms::new(client, resource.into())); + let settings = KeyManagementSettings::default(); + let value = get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| { + Some(STANDARD.encode("encrypted")) + }) + .await + .unwrap() + .unwrap(); + assert_eq!(value.as_str(), Some(" value\n")); + assert!(matches!( + get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| Some(format!( + " {}", + STANDARD.encode("encrypted") + ))) + .await, + Err(Error::InvalidCiphertext) + )); + assert!(matches!( + get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| None).await, + Err(Error::MissingCiphertext) + )); +} diff --git a/litellm-rust/crates/secrets/tests/handler.rs b/litellm-rust/crates/secrets/tests/handler.rs deleted file mode 100644 index 2a8b7070522..00000000000 --- a/litellm-rust/crates/secrets/tests/handler.rs +++ /dev/null @@ -1,363 +0,0 @@ -#[cfg(feature = "aws")] -#[tokio::test] -async fn aws_handler_reads_ciphertext_decodes_trims_and_redacts() { - use aws_sdk_kms::{ - Client, - config::{BehaviorVersion, Credentials, Region}, - }; - use base64::{Engine, engine::general_purpose::STANDARD}; - use litellm_secrets::{ - Error, KeyManagementSettings, SecretManager, aws::AwsKms, get_secret_from_manager, - }; - use wiremock::{Mock, MockServer, ResponseTemplate, matchers::body_json}; - - let server = MockServer::start().await; - Mock::given(body_json( - serde_json::json!({"CiphertextBlob": STANDARD.encode("encrypted")}), - )) - .respond_with( - ResponseTemplate::new(200) - .set_body_json(serde_json::json!({"Plaintext":STANDARD.encode(" value\n")})), - ) - .expect(1) - .mount(&server) - .await; - let client = Client::from_conf( - aws_sdk_kms::Config::builder() - .behavior_version(BehaviorVersion::latest()) - .region(Region::new("us-east-1")) - .credentials_provider(Credentials::new("test", "test", None, None, "test")) - .endpoint_url(server.uri()) - .build(), - ); - let manager = SecretManager::AwsKms(AwsKms::new(client)); - let settings = KeyManagementSettings::default(); - let value = get_secret_from_manager(&manager, "KEY", &settings, &|name: &str| { - assert_eq!(name, "KEY"); - Some(format!(" {}\n", STANDARD.encode("encrypted"))) - }) - .await - .unwrap() - .unwrap(); - assert_eq!(value.as_str(), Some("value")); - assert!(!format!("{value:?}").contains("value")); - assert!(matches!( - get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| None).await, - Err(Error::MissingCiphertext) - )); - assert!(matches!( - get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| Some("abc".into())).await, - Err(Error::InvalidCiphertext) - )); -} - -#[cfg(feature = "google")] -#[tokio::test] -async fn google_handler_requires_canonical_base64_and_preserves_plaintext_whitespace() { - use base64::{Engine, engine::general_purpose::STANDARD}; - use google_cloud_kms_v1::client::KeyManagementService; - use litellm_secrets::{ - Error, KeyManagementSettings, SecretManager, get_secret_from_manager, google::GoogleKms, - }; - use wiremock::{ - Mock, MockServer, ResponseTemplate, - matchers::{body_json, path}, - }; - - let server = MockServer::start().await; - let resource = "projects/project/locations/global/keyRings/ring/cryptoKeys/key"; - Mock::given(path(format!("/v1/{resource}:decrypt"))) - .and(body_json( - serde_json::json!({"ciphertext":STANDARD.encode("encrypted")}), - )) - .respond_with( - ResponseTemplate::new(200) - .set_body_json(serde_json::json!({"plaintext":STANDARD.encode(" value\n")})), - ) - .expect(1) - .mount(&server) - .await; - let client = KeyManagementService::builder() - .with_endpoint(server.uri()) - .with_credentials(google_cloud_auth::credentials::anonymous::Builder::new().build()) - .build() - .await - .unwrap(); - let manager = SecretManager::GoogleKms(GoogleKms::new(client, resource.into())); - let settings = KeyManagementSettings::default(); - let value = get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| { - Some(STANDARD.encode("encrypted")) - }) - .await - .unwrap() - .unwrap(); - assert_eq!(value.as_str(), Some(" value\n")); - assert!(matches!( - get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| Some(format!( - " {}", - STANDARD.encode("encrypted") - ))) - .await, - Err(Error::InvalidCiphertext) - )); - assert!(matches!( - get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| None).await, - Err(Error::MissingCiphertext) - )); -} -#[cfg(feature = "hashicorp")] -#[tokio::test] -async fn hashicorp_handler_resolves_found_missing_and_failed_values() { - use std::sync::Arc; - - use litellm_core_utils::settings::Lookup; - use litellm_secrets::{ - Error, FailurePolicy, KeyManagementSettings, SecretManager, SecretManagerState, - SecretResolver, hashicorp::HashicorpVault, hashicorp::HashicorpVaultConfig, - }; - use wiremock::{ - Mock, MockServer, ResponseTemplate, - matchers::{method, path}, - }; - - let found_server = MockServer::start().await; - Mock::given(method("GET")) - .and(path("/v1/secret/data/KEY")) - .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ - "data": { - "data": {"key": "remote"}, - "metadata": { - "created_time": "", - "deletion_time": "", - "custom_metadata": null, - "destroyed": false, - "version": 1 - } - }, - "lease_id": "", - "lease_duration": 0, - "renewable": false, - "request_id": "", - "warnings": null, - "wrap_info": null - }))) - .mount(&found_server) - .await; - let found_environment: Arc = Arc::new({ - let address = found_server.uri(); - move |name: &str| match name { - "HCP_VAULT_ADDR" => Some(address.clone()), - "HCP_VAULT_TOKEN" => Some("token".into()), - _ => None, - } - }); - let found_config = HashicorpVaultConfig::from_environment(found_environment.as_ref()).unwrap(); - let found_manager = HashicorpVault::from_config(found_config, true).unwrap(); - let found_resolver = SecretResolver::new( - Arc::new(SecretManagerState::new( - SecretManager::HashicorpVault(found_manager), - KeyManagementSettings { - hosted_keys: Some(vec!["KEY".into()]), - ..Default::default() - }, - )), - Arc::new(|_: &str| None), - litellm_secrets::OidcResolver::default(), - ); - assert_eq!( - found_resolver - .get_secret_str("KEY", None) - .await - .unwrap() - .unwrap() - .expose(), - "remote" - ); - - let missing_server = MockServer::start().await; - Mock::given(method("GET")) - .respond_with( - ResponseTemplate::new(404).set_body_json(serde_json::json!({"errors": ["missing"]})), - ) - .mount(&missing_server) - .await; - let missing_environment: Arc = Arc::new({ - let address = missing_server.uri(); - move |name: &str| match name { - "HCP_VAULT_ADDR" => Some(address.clone()), - "HCP_VAULT_TOKEN" => Some("token".into()), - _ => None, - } - }); - let missing_config = - HashicorpVaultConfig::from_environment(missing_environment.as_ref()).unwrap(); - let missing_manager = HashicorpVault::from_config(missing_config, true).unwrap(); - let missing_state = SecretManagerState::new( - SecretManager::HashicorpVault(missing_manager), - KeyManagementSettings { - hosted_keys: Some(vec!["KEY".into()]), - ..Default::default() - }, - ); - let missing = litellm_secrets::get_secret_from_manager( - missing_state.backend().unwrap(), - "KEY", - missing_state.settings().unwrap(), - &|_: &str| None, - ) - .await - .unwrap(); - assert!(missing.is_none()); - - let failed_server = MockServer::start().await; - Mock::given(method("GET")) - .respond_with( - ResponseTemplate::new(500).set_body_json(serde_json::json!({"errors": ["failed"]})), - ) - .mount(&failed_server) - .await; - let failed_environment: Arc = Arc::new({ - let address = failed_server.uri(); - move |name: &str| match name { - "HCP_VAULT_ADDR" => Some(address.clone()), - "HCP_VAULT_TOKEN" => Some("token".into()), - _ => None, - } - }); - let failed_config = - HashicorpVaultConfig::from_environment(failed_environment.as_ref()).unwrap(); - let failed_manager = HashicorpVault::from_config(failed_config, true).unwrap(); - let failed_state = SecretManagerState::new( - SecretManager::HashicorpVault(failed_manager), - KeyManagementSettings { - hosted_keys: Some(vec!["KEY".into()]), - ..Default::default() - }, - ); - let failed_resolver = SecretResolver::new( - Arc::new(failed_state), - Arc::new(|_: &str| None), - litellm_secrets::OidcResolver::default(), - ) - .with_failure_policy(FailurePolicy::Propagate); - assert!(matches!( - failed_resolver.get_secret_str("KEY", None).await, - Err(Error::Hashicorp( - litellm_secrets::hashicorp::Error::Status { status: 500 } - )) - )); -} - -#[cfg(feature = "azure")] -#[tokio::test] -async fn azure_handler_reads_missing_and_failed_secrets() { - use litellm_secrets::{ - Error, KeyManagementSettings, KeyManagementSystem, SecretManager, azure::AzureKeyVault, - get_secret_from_manager, - }; - use wiremock::{ - Mock, MockServer, ResponseTemplate, - matchers::{path, query_param}, - }; - - let server = MockServer::start().await; - Mock::given(path("/secrets/KEY")) - .and(query_param("api-version", "7.4")) - .respond_with( - ResponseTemplate::new(200).set_body_json(serde_json::json!({"value": "value"})), - ) - .expect(1) - .mount(&server) - .await; - let manager = SecretManager::AzureKeyVault( - AzureKeyVault::with_client( - reqwest::Client::new(), - server.uri().parse().unwrap(), - std::sync::Arc::new(|name: &str| (name == "AZURE_AD_TOKEN").then(|| "fake".to_owned())), - ) - .unwrap(), - ); - assert_eq!(manager.system(), KeyManagementSystem::AzureKeyVault); - let settings = KeyManagementSettings::default(); - let value = get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| None) - .await - .unwrap() - .unwrap(); - assert_eq!(value.as_str(), Some("value")); - - let not_found = Mock::given(path("/secrets/MISSING")) - .respond_with(ResponseTemplate::new(404)) - .expect(1) - .mount_as_scoped(&server) - .await; - assert_eq!( - get_secret_from_manager(&manager, "MISSING", &settings, &|_: &str| None) - .await - .unwrap(), - None - ); - drop(not_found); - - Mock::given(path("/secrets/FAILED")) - .respond_with(ResponseTemplate::new(500)) - .expect(1) - .mount(&server) - .await; - assert!(matches!( - get_secret_from_manager(&manager, "FAILED", &settings, &|_: &str| None).await, - Err(Error::Azure(_)) - )); -} - -#[cfg(feature = "cyberark")] -#[tokio::test] -async fn cyberark_handler_reads_values_and_surfaces_errors() { - use std::time::Duration; - - use litellm_secrets::{ - Error, KeyManagementSettings, SecretManager, SecretValue, cyberark::CyberArkSecretManager, - get_secret_from_manager, - }; - use wiremock::{ - Mock, MockServer, ResponseTemplate, - matchers::{body_string, path}, - }; - - let server = MockServer::start().await; - Mock::given(path("/authn/acct/admin/authenticate")) - .and(body_string("k3y")) - .respond_with(ResponseTemplate::new(200).set_body_string("token")) - .mount(&server) - .await; - Mock::given(path("/secrets/acct/variable/KEY")) - .respond_with(ResponseTemplate::new(200).set_body_string("value")) - .mount(&server) - .await; - let manager = SecretManager::Cyberark(CyberArkSecretManager::with_client( - reqwest::Client::new(), - server.uri().parse().unwrap(), - "acct".into(), - "admin".into(), - SecretValue::new("k3y"), - Some(Duration::from_secs(60)), - )); - assert_eq!( - manager.system(), - litellm_secrets::KeyManagementSystem::Cyberark - ); - let settings = KeyManagementSettings::default(); - let value = get_secret_from_manager(&manager, "KEY", &settings, &|_: &str| None) - .await - .unwrap() - .unwrap(); - assert_eq!(value.as_str(), Some("value")); - - Mock::given(path("/secrets/acct/variable/ERROR")) - .respond_with(ResponseTemplate::new(500)) - .mount(&server) - .await; - assert!(matches!( - get_secret_from_manager(&manager, "ERROR", &settings, &|_: &str| None).await, - Err(Error::Cyberark(_)) - )); -} diff --git a/litellm-rust/crates/secrets/tests/hashicorp.rs b/litellm-rust/crates/secrets/tests/hashicorp.rs new file mode 100644 index 00000000000..bc35b88018e --- /dev/null +++ b/litellm-rust/crates/secrets/tests/hashicorp.rs @@ -0,0 +1,143 @@ +#![cfg(feature = "hashicorp")] + +#[tokio::test] +async fn hashicorp_handler_resolves_found_missing_and_failed_values() { + use std::sync::Arc; + + use litellm_core_utils::settings::Lookup; + use litellm_secrets::{ + Error, FailurePolicy, KeyManagementSettings, SecretManager, SecretManagerState, + SecretResolver, hashicorp::HashicorpVault, hashicorp::HashicorpVaultConfig, + }; + use wiremock::{ + Mock, MockServer, ResponseTemplate, + matchers::{method, path}, + }; + + let found_server = MockServer::start().await; + Mock::given(method("GET")) + .and(path("/v1/secret/data/KEY")) + .respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({ + "data": { + "data": {"key": "remote"}, + "metadata": { + "created_time": "", + "deletion_time": "", + "custom_metadata": null, + "destroyed": false, + "version": 1 + } + }, + "lease_id": "", + "lease_duration": 0, + "renewable": false, + "request_id": "", + "warnings": null, + "wrap_info": null + }))) + .mount(&found_server) + .await; + let found_environment: Arc = Arc::new({ + let address = found_server.uri(); + move |name: &str| match name { + "HCP_VAULT_ADDR" => Some(address.clone()), + "HCP_VAULT_TOKEN" => Some("token".into()), + _ => None, + } + }); + let found_config = HashicorpVaultConfig::from_environment(found_environment.as_ref()).unwrap(); + let found_manager = HashicorpVault::from_config(found_config, true).unwrap(); + let found_resolver = SecretResolver::new_python_compatible( + Arc::new(SecretManagerState::new( + SecretManager::HashicorpVault(found_manager), + KeyManagementSettings { + hosted_keys: Some(vec!["KEY".into()]), + ..Default::default() + }, + )), + Arc::new(|_: &str| None), + litellm_secrets::OidcResolver::default(), + ); + assert_eq!( + found_resolver + .get_secret_str("KEY", None) + .await + .unwrap() + .unwrap() + .expose(), + "remote" + ); + + let missing_server = MockServer::start().await; + Mock::given(method("GET")) + .respond_with( + ResponseTemplate::new(404).set_body_json(serde_json::json!({"errors": ["missing"]})), + ) + .mount(&missing_server) + .await; + let missing_environment: Arc = Arc::new({ + let address = missing_server.uri(); + move |name: &str| match name { + "HCP_VAULT_ADDR" => Some(address.clone()), + "HCP_VAULT_TOKEN" => Some("token".into()), + _ => None, + } + }); + let missing_config = + HashicorpVaultConfig::from_environment(missing_environment.as_ref()).unwrap(); + let missing_manager = HashicorpVault::from_config(missing_config, true).unwrap(); + let missing_state = SecretManagerState::new( + SecretManager::HashicorpVault(missing_manager), + KeyManagementSettings { + hosted_keys: Some(vec!["KEY".into()]), + ..Default::default() + }, + ); + let missing = litellm_secrets::get_secret_from_manager( + missing_state.backend().unwrap(), + "KEY", + missing_state.settings().unwrap(), + &|_: &str| None, + ) + .await + .unwrap(); + assert!(missing.is_none()); + + let failed_server = MockServer::start().await; + Mock::given(method("GET")) + .respond_with( + ResponseTemplate::new(500).set_body_json(serde_json::json!({"errors": ["failed"]})), + ) + .mount(&failed_server) + .await; + let failed_environment: Arc = Arc::new({ + let address = failed_server.uri(); + move |name: &str| match name { + "HCP_VAULT_ADDR" => Some(address.clone()), + "HCP_VAULT_TOKEN" => Some("token".into()), + _ => None, + } + }); + let failed_config = + HashicorpVaultConfig::from_environment(failed_environment.as_ref()).unwrap(); + let failed_manager = HashicorpVault::from_config(failed_config, true).unwrap(); + let failed_state = SecretManagerState::new( + SecretManager::HashicorpVault(failed_manager), + KeyManagementSettings { + hosted_keys: Some(vec!["KEY".into()]), + ..Default::default() + }, + ); + let failed_resolver = SecretResolver::new_python_compatible( + Arc::new(failed_state), + Arc::new(|_: &str| None), + litellm_secrets::OidcResolver::default(), + ) + .with_failure_policy(FailurePolicy::Propagate); + assert!(matches!( + failed_resolver.get_secret_str("KEY", None).await, + Err(Error::Hashicorp( + litellm_secrets::hashicorp::Error::Status { status: 500 } + )) + )); +} diff --git a/litellm-rust/crates/secrets/tests/oidc.rs b/litellm-rust/crates/secrets/tests/oidc.rs index b17e7de7f9d..afc49e8231d 100644 --- a/litellm-rust/crates/secrets/tests/oidc.rs +++ b/litellm-rust/crates/secrets/tests/oidc.rs @@ -185,6 +185,8 @@ async fn file_allowlist_resolves_symlinks_while_environment_paths_remain_explici #[case::string_expiry(serde_json::json!("999"), 2)] #[case::fractional_expiry(serde_json::json!(1060.9), 2)] #[case::negative_expiry(serde_json::json!(-1), 2)] +#[case::boolean_expiry(serde_json::json!(true), 2)] +#[case::padded_numeric_expiry(serde_json::json!(" 999 "), 2)] #[case::null_expiry(serde_json::Value::Null, 1)] #[case::unreadable_expiry(serde_json::json!("invalid"), 1)] #[case::nonfinite_expiry(serde_json::json!("NaN"), 1)] @@ -239,8 +241,9 @@ async fn google_oidc_requires_its_build_feature() { )); } +#[cfg(not(feature = "azure"))] #[tokio::test] -async fn azure_oidc_without_a_token_file_requires_an_unimplemented_backend() { +async fn azure_oidc_without_a_token_file_requires_its_build_feature() { assert!(matches!( OidcResolver::default() .resolve("oidc/azure/scope", environment(&[]).as_ref()) @@ -293,3 +296,193 @@ async fn unreadable_expiry_keeps_python_cache_fallback(#[case] token: &str) { ); } } + +#[cfg(feature = "azure")] +#[rstest::rstest] +#[case::success(false)] +#[case::failed(true)] +#[tokio::test] +async fn azure_oidc_acquires_the_requested_scope_and_preserves_failures(#[case] failed: bool) { + struct Provider(bool); + impl litellm_secrets::azure::AzureTokenProvider for Provider { + fn get_token<'a>( + &'a self, + scope: &'a str, + environment: &'a (dyn Lookup + Send + Sync), + ) -> std::pin::Pin< + Box< + dyn std::future::Future< + Output = Result< + litellm_secrets::SecretValue, + litellm_secrets::azure::Error, + >, + > + Send + + 'a, + >, + > { + Box::pin(async move { + assert_eq!(scope, "api://audience/path"); + assert_eq!( + environment.get("AZURE_CLIENT_ID").as_deref(), + Some("client-id") + ); + if self.0 { + Err(litellm_secrets::azure::Error::MissingCredentials) + } else { + Ok(litellm_secrets::SecretValue::new("azure-token")) + } + }) + } + } + let oidc = OidcResolver::default().with_azure_token_provider(Arc::new(Provider(failed))); + let resolver = SecretResolver::new_python_compatible( + Arc::new(SecretManagerState::default()), + environment(&[("AZURE_CLIENT_ID", "client-id")]), + oidc, + ); + let result = resolver + .get_secret_str( + "oidc/azure/api://audience/path", + Some(litellm_secrets::SecretValue::new("fallback")), + ) + .await; + if failed { + assert!(matches!(result, Err(Error::Azure(_)))); + } else { + assert_eq!(result.unwrap().unwrap().expose(), "azure-token"); + } +} + +#[rstest::rstest] +#[case::circleci("oidc/circleci/audience")] +#[case::circleci_v2("oidc/circleci_v2/audience")] +#[case::env("oidc/env/MISSING")] +#[case::env_path("oidc/env_path/MISSING")] +#[tokio::test] +async fn missing_oidc_environment_is_an_error(#[case] reference: &str) { + assert!(matches!( + OidcResolver::default() + .resolve(reference, environment(&[]).as_ref()) + .await, + Err(Error::MissingEnvironment) + )); +} + +#[cfg(feature = "google")] +#[tokio::test] +async fn google_oidc_failures_are_not_cached_or_hidden_by_defaults() { + let server = MockServer::start().await; + Mock::given(method("GET")) + .respond_with(ResponseTemplate::new(403)) + .expect(2) + .mount(&server) + .await; + let resolver = SecretResolver::new_python_compatible( + Arc::new(SecretManagerState::default()), + environment(&[]), + OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()), + ); + for _ in 0..2 { + assert!(matches!( + resolver + .get_secret("oidc/google/audience", Some(Secret::Bool(true))) + .await, + Err(Error::OidcStatus(403)) + )); + } +} + +#[cfg(feature = "google")] +#[rstest::rstest] +#[case::short_lived(Some(1180), 120)] +#[case::long_lived(Some(100000), 3540)] +#[case::opaque(None, 3540)] +#[tokio::test] +async fn google_tokens_expire_at_the_python_cache_deadline( + #[case] expiry: Option, + #[case] ttl: u64, +) { + use base64::{Engine, engine::general_purpose::URL_SAFE_NO_PAD}; + use std::time::{Duration, UNIX_EPOCH}; + let server = MockServer::start().await; + let token = expiry.map_or_else( + || "opaque-token".to_owned(), + |expiry| { + format!( + "{}.{}.signature", + URL_SAFE_NO_PAD.encode(r#"{"alg":"RS256"}"#), + URL_SAFE_NO_PAD.encode(serde_json::json!({"exp":expiry}).to_string()) + ) + }, + ); + Mock::given(method("GET")) + .respond_with(ResponseTemplate::new(200).set_body_string(&token)) + .expect(2) + .mount(&server) + .await; + let resolver = OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()) + .with_clock(|| UNIX_EPOCH + Duration::from_secs(1000)); + assert_eq!( + resolver + .resolve("oidc/google/audience", environment(&[]).as_ref()) + .await + .unwrap() + .unwrap() + .expose(), + token + ); + let before_deadline = match ttl { + 120 => resolver.with_clock(|| UNIX_EPOCH + Duration::from_secs(1119)), + 3540 => resolver.with_clock(|| UNIX_EPOCH + Duration::from_secs(4539)), + _ => unreachable!(), + }; + assert_eq!( + before_deadline + .resolve("oidc/google/audience", environment(&[]).as_ref()) + .await + .unwrap() + .unwrap() + .expose(), + token + ); + let at_deadline = match ttl { + 120 => before_deadline.with_clock(|| UNIX_EPOCH + Duration::from_secs(1120)), + 3540 => before_deadline.with_clock(|| UNIX_EPOCH + Duration::from_secs(4540)), + _ => unreachable!(), + }; + assert_eq!( + at_deadline + .resolve("oidc/google/audience", environment(&[]).as_ref()) + .await + .unwrap() + .unwrap() + .expose(), + token + ); +} + +#[cfg(feature = "google")] +#[tokio::test] +async fn google_cache_uses_payload_expiry_without_requiring_a_jwt_header() { + use std::time::{Duration, UNIX_EPOCH}; + let server = MockServer::start().await; + let token = "ignored.eyJleHAiOjF9.ignored"; + Mock::given(method("GET")) + .respond_with(ResponseTemplate::new(200).set_body_string(token)) + .expect(2) + .mount(&server) + .await; + let resolver = OidcResolver::new(reqwest::Client::new(), server.uri().parse().unwrap()) + .with_clock(|| UNIX_EPOCH + Duration::from_secs(1000)); + for _ in 0..2 { + assert_eq!( + resolver + .resolve("oidc/google/audience", environment(&[]).as_ref()) + .await + .unwrap() + .unwrap() + .expose(), + token + ); + } +} diff --git a/litellm-rust/crates/secrets/tests/resolution.rs b/litellm-rust/crates/secrets/tests/resolution.rs index 98b34a2c72e..bed762adc59 100644 --- a/litellm-rust/crates/secrets/tests/resolution.rs +++ b/litellm-rust/crates/secrets/tests/resolution.rs @@ -1,13 +1,15 @@ -use std::sync::Arc; +use std::{future::Future, pin::Pin, sync::Arc}; +use litellm_core_utils::settings::Lookup; use litellm_secrets::{ - Error, OidcResolver, Secret, SecretManagerState, SecretResolver, SecretValue, + Error, ExternalSecretManager, FailurePolicy, KeyManagementSettings, KeyManagementSystem, + OidcResolver, Secret, SecretManager, SecretManagerState, SecretResolver, SecretValue, secret_manager_would_be_consulted, }; fn resolver(value: Option<&str>) -> SecretResolver { let value = value.map(str::to_owned); - SecretResolver::new( + SecretResolver::new_python_compatible( Arc::new(SecretManagerState::default()), Arc::new(move |_: &str| value.clone()), OidcResolver::default(), @@ -15,76 +17,131 @@ fn resolver(value: Option<&str>) -> SecretResolver { } #[rstest::rstest] -#[case("true", Some(true))] -#[case(" FALSE ", Some(false))] -#[case("(True)", None)] -#[case("False # comment", None)] -#[case("1", None)] -#[case("secret", None)] +#[case::environment(false)] +#[case::manager(true)] #[tokio::test] -async fn conversion_is_explicit_and_independent_of_manager_configuration( +async fn native_reads_preserve_strings_and_report_conversion_errors(#[case] managed: bool) { + for raw in ["True", " FALSE ", "", "(True)", "{\"key\":1}"] { + let state = if managed { + SecretManagerState::new( + SecretManager::External(Arc::new(FixedManager::custom(Ok(Some(Secret::String( + SecretValue::new(raw), + )))))), + KeyManagementSettings::default(), + ) + } else { + SecretManagerState::default() + }; + let resolver = SecretResolver::new( + Arc::new(state), + Arc::new(move |_: &str| Some(raw.to_owned())), + OidcResolver::default(), + ); + assert_eq!( + resolver + .get_secret_str("key", None) + .await + .unwrap() + .unwrap() + .expose(), + raw + ); + assert_eq!( + resolver.get_secret("key", None).await.unwrap(), + Some(Secret::String(SecretValue::new(raw))) + ); + if raw == "True" || raw == " FALSE " { + assert_eq!( + resolver.get_secret_bool("key", None).await.unwrap(), + Some(raw == "True") + ); + } else { + assert!(matches!( + resolver.get_secret_bool("key", None).await, + Err(Error::TypeMismatch { + expected: "boolean" + }) + )); + } + } +} + +#[tokio::test] +async fn native_defaults_apply_to_absence_but_never_hide_provider_failures() { + for reply in [Ok(None), Err(())] { + let resolver = SecretResolver::new( + Arc::new(SecretManagerState::new( + SecretManager::External(Arc::new(FixedManager::custom(reply.clone()))), + KeyManagementSettings::default(), + )), + Arc::new(|_: &str| None), + OidcResolver::default(), + ); + let result = resolver + .get_secret_str("key", Some(SecretValue::new("default"))) + .await; + match reply { + Ok(None) => assert_eq!(result.unwrap().unwrap().expose(), "default"), + Err(()) => assert!(matches!(result, Err(Error::MissingCiphertext))), + Ok(Some(_)) => unreachable!(), + } + } +} + +#[rstest::rstest] +#[case::lowercase_true("true", Some(true))] +#[case::padded_false(" FALSE ", Some(false))] +#[case::capitalized_true("True", Some(true))] +#[case::parenthesized("(True)", None)] +#[case::commented("False # comment", None)] +#[case::number("1", None)] +#[case::text("secret", None)] +#[tokio::test] +async fn environment_values_are_coerced_like_str_to_bool( #[case] input: &str, #[case] boolean: Option, ) { let resolver = resolver(Some(input)); assert_eq!( resolver.get_secret("key", None).await.unwrap(), - Some(Secret::String(SecretValue::new(input))) + Some(boolean.map_or_else(|| Secret::String(SecretValue::new(input)), Secret::Bool)) ); assert_eq!( resolver .get_secret_str("key", None) .await .unwrap() - .unwrap() - .expose(), - input + .as_ref() + .map(SecretValue::expose), + boolean.is_none().then_some(input) + ); + assert_eq!( + resolver.get_secret_bool("key", Some(true)).await.unwrap(), + boolean ); - match boolean { - Some(value) => assert_eq!( - resolver.get_secret_bool("key", None).await.unwrap(), - Some(value) - ), - None => assert!(matches!( - resolver.get_secret_bool("key", Some(true)).await, - Err(Error::TypeMismatch { - expected: "boolean" - }) - )), - } } -#[rstest::rstest] #[tokio::test] -async fn defaults_apply_only_to_absence() { +async fn defaults_never_replace_an_absent_secret() { let missing = resolver(None); - assert_eq!(missing.get_secret("key", None).await.unwrap(), None); + assert_eq!( + missing + .get_secret("key", Some(Secret::Bool(false))) + .await + .unwrap(), + None + ); assert_eq!( missing.get_secret_bool("key", Some(false)).await.unwrap(), - Some(false) + None ); assert_eq!( missing .get_secret_str("key", Some(SecretValue::new("default"))) .await - .unwrap() - .unwrap() - .expose(), - "default" + .unwrap(), + None ); - for value in [ - Secret::Bool(false), - Secret::from_json(serde_json::json!({"key":1})), - Secret::from_json(serde_json::Value::Null), - ] { - assert_eq!( - missing - .get_secret("key", Some(value.clone())) - .await - .unwrap(), - Some(value) - ); - } assert_eq!( resolver(Some("")) .get_secret_str("key", Some(SecretValue::new("default"))) @@ -96,6 +153,124 @@ async fn defaults_apply_only_to_absence() { ); } +struct FixedManager { + reply: Result, ()>, + system: KeyManagementSystem, +} + +impl FixedManager { + fn custom(reply: Result, ()>) -> Self { + Self { + reply, + system: KeyManagementSystem::Custom, + } + } +} + +impl ExternalSecretManager for FixedManager { + fn system(&self) -> KeyManagementSystem { + self.system + } + + fn read_secret<'a>( + &'a self, + _name: &'a str, + _settings: &'a KeyManagementSettings, + _environment: &'a (dyn Lookup + Send + Sync), + ) -> Pin, Error>> + Send + 'a>> { + Box::pin(async move { self.reply.clone().map_err(|()| Error::MissingCiphertext) }) + } +} + +fn managed(reply: Result, ()>, environment: Option<&'static str>) -> SecretResolver { + SecretResolver::new_python_compatible( + Arc::new(SecretManagerState::new( + SecretManager::External(Arc::new(FixedManager::custom(reply))), + KeyManagementSettings::default(), + )), + Arc::new(move |_: &str| environment.map(str::to_owned)), + OidcResolver::default(), + ) + .with_failure_policy(FailurePolicy::EnvironmentFallback) +} + +#[tokio::test] +async fn custom_manager_absence_uses_environment_instead_of_the_default() { + assert_eq!( + managed(Ok(None), Some("environment")) + .get_secret("key", Some(Secret::Bool(true))) + .await + .unwrap(), + Some(Secret::String(SecretValue::new("environment"))) + ); +} + +#[rstest::rstest] +#[case::capitalized_true("True", Some(Secret::Bool(true)), Some(true))] +#[case::parenthesized_false("(False)", Some(Secret::Bool(false)), Some(false))] +#[case::lowercase_true("true", None, Some(true))] +#[case::number("1", None, None)] +#[case::text("secret", None, None)] +#[tokio::test] +async fn manager_strings_are_coerced_like_literal_eval( + #[case] input: &'static str, + #[case] literal: Option, + #[case] boolean: Option, +) { + let resolver = managed(Ok(Some(Secret::String(SecretValue::new(input)))), None); + assert_eq!( + resolver.get_secret("key", None).await.unwrap(), + Some( + literal + .clone() + .unwrap_or_else(|| Secret::String(SecretValue::new(input))) + ) + ); + assert_eq!( + resolver + .get_secret_str("key", None) + .await + .unwrap() + .as_ref() + .map(SecretValue::expose), + literal.is_none().then_some(input) + ); + assert_eq!( + resolver.get_secret_bool("key", None).await.unwrap(), + boolean + ); +} + +#[rstest::rstest] +#[case::boolean(Secret::Bool(false))] +#[case::object(Secret::from_json(serde_json::json!({"key": 1})))] +#[case::null(Secret::from_json(serde_json::Value::Null))] +#[tokio::test] +async fn non_string_manager_values_resolve_to_none(#[case] value: Secret) { + let resolver = managed(Ok(Some(value)), Some("environment")); + assert_eq!(resolver.get_secret("key", None).await.unwrap(), None); + assert_eq!(resolver.get_secret_str("key", None).await.unwrap(), None); + assert_eq!(resolver.get_secret_bool("key", None).await.unwrap(), None); +} + +#[rstest::rstest] +#[case::capitalized_true(Some("True"), Some(Secret::Bool(true)))] +#[case::lowercase_true(Some("true"), Some(Secret::String(SecretValue::new("true"))))] +#[case::missing(None, None)] +#[tokio::test] +async fn manager_failures_fall_back_to_the_environment_like_literal_eval( + #[case] environment: Option<&'static str>, + #[case] expected: Option, +) { + assert_eq!( + managed(Err(()), environment) + .get_secret("key", Some(Secret::Bool(false))) + .await + .unwrap(), + expected + ); +} + #[tokio::test] async fn prefix_is_removed_once_and_resolved_from_environment() { let state = SecretManagerState::default(); @@ -103,7 +278,7 @@ async fn prefix_is_removed_once_and_resolved_from_environment() { &state, "os.environ/os.environ/KEY" )); - let resolver = SecretResolver::new( + let resolver = SecretResolver::new_python_compatible( Arc::new(state), Arc::new(|name: &str| (name == "os.environ/KEY").then(|| "value".into())), OidcResolver::default(), @@ -129,234 +304,82 @@ async fn resolver_future_can_run_on_a_tokio_worker() { assert_eq!(result.unwrap().expose(), "worker-value"); } -#[cfg(feature = "aws")] -mod aws { - use super::*; - use litellm_secrets::{ - AccessMode, FailurePolicy, KeyManagementSettings, SecretManager, aws::AwsSecretsManagerV2, - }; - use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method}; +#[rstest::rstest] +#[case::missing(None, None)] +#[case::empty(Some(""), None)] +#[case::whitespace(Some(" \t\n"), None)] +#[case::text(Some("abc"), Some("abc"))] +#[case::padded(Some(" xyz "), Some("xyz"))] +#[case::python_controls(Some("\u{1c}\u{1d}\u{1e}\u{1f}"), None)] +#[case::unicode(Some("\u{a0}Ï€\u{2003}"), Some("Ï€"))] +fn normalization_matches_python_without_changing_embedded_whitespace( + #[case] input: Option<&str>, + #[case] expected: Option<&str>, +) { + assert_eq!( + litellm_secrets::normalize_nonempty_secret_str(input), + expected + ); +} - fn state(server: &MockServer, settings: KeyManagementSettings) -> SecretManagerState { - let endpoint = server.uri(); - let environment = Arc::new(move |name: &str| match name { - "AWS_REGION_NAME" => Some("us-east-1".into()), - "AWS_ACCESS_KEY_ID" | "AWS_SECRET_ACCESS_KEY" => Some("test".into()), - "AWS_BEDROCK_RUNTIME_ENDPOINT" => Some(endpoint.clone()), - _ => None, - }); - let manager = - AwsSecretsManagerV2::load_aws_secret_manager(Some(true), settings.clone(), environment) - .unwrap() - .unwrap(); - SecretManagerState::new(SecretManager::AwsSecretsManagerV2(manager), settings) - } - - #[rstest::rstest] - #[case::missing(400, serde_json::json!({"__type":"ResourceNotFoundException"}), false)] - #[case::denied(400, serde_json::json!({"__type":"AccessDeniedException"}), true)] - #[case::malformed(200, serde_json::json!({}), true)] - #[tokio::test] - async fn failure_policy_preserves_errors_and_fallback_precedence( - #[case] status: u16, - #[case] body: serde_json::Value, - #[case] fails: bool, - #[values(FailurePolicy::Propagate, FailurePolicy::EnvironmentFallback)] - policy: FailurePolicy, - #[values(None, Some("environment"))] environment: Option<&'static str>, - #[values(None, Some("default"))] default: Option<&str>, - ) { - let server = MockServer::start().await; - Mock::given(method("POST")) - .respond_with(ResponseTemplate::new(status).set_body_json(body)) - .expect(1) - .mount(&server) - .await; - let resolver = SecretResolver::new( - Arc::new(state(&server, KeyManagementSettings::default())), - Arc::new(move |_: &str| environment.map(str::to_owned)), - OidcResolver::default(), - ) - .with_failure_policy(policy); - let result = resolver - .get_secret_str("KEY", default.map(SecretValue::new)) - .await; - let fallback = environment.or(default); - if fails && (policy == FailurePolicy::Propagate || fallback.is_none()) { - assert!(matches!(result, Err(Error::Aws(_)))); - } else { - assert_eq!(result.unwrap().as_ref().map(SecretValue::expose), fallback); - } - } - - #[rstest::rstest] - #[case::boolean(serde_json::json!(false))] - #[case::object(serde_json::json!({"key":1}))] - #[case::null(serde_json::Value::Null)] - #[case::string(serde_json::json!("true"))] - #[tokio::test] - async fn typed_values_survive_resolution_and_accessors_reject_wrong_types( - #[case] value: serde_json::Value, - ) { - let server = MockServer::start().await; - Mock::given(method("POST")) - .respond_with(ResponseTemplate::new(200).set_body_json( - serde_json::json!({"SecretString":serde_json::json!({"KEY":value}).to_string()}), - )) - .expect(3) - .mount(&server) - .await; - let settings = KeyManagementSettings { - primary_secret_name: Some("primary".into()), - ..Default::default() - }; - let resolver = SecretResolver::new( - Arc::new(state(&server, settings)), - Arc::new(|_: &str| Some("fallback".into())), - OidcResolver::default(), - ); - assert_eq!( - resolver - .get_secret("KEY", Some(Secret::Bool(true))) - .await - .unwrap(), - Some(Secret::from_json(value.clone())) - ); - match &value { - serde_json::Value::String(text) => assert_eq!( - resolver - .get_secret_str("KEY", None) - .await - .unwrap() - .unwrap() - .expose(), - text - ), - _ => assert!(matches!( - resolver.get_secret_str("KEY", None).await, - Err(Error::TypeMismatch { expected: "string" }) - )), - } - match value { - serde_json::Value::Bool(boolean) => assert_eq!( - resolver.get_secret_bool("KEY", None).await.unwrap(), - Some(boolean) - ), - serde_json::Value::String(_) => assert_eq!( - resolver.get_secret_bool("KEY", None).await.unwrap(), - Some(true) - ), - _ => assert!(matches!( - resolver.get_secret_bool("KEY", None).await, - Err(Error::TypeMismatch { - expected: "boolean" - }) - )), - } - } - - #[rstest::rstest] - #[tokio::test] - async fn gating_prediction_matches_actual_lookup( - #[values(AccessMode::ReadOnly, AccessMode::WriteOnly, AccessMode::ReadAndWrite)] - access_mode: AccessMode, - #[values(None, Some(vec![]), Some(vec!["KEY".into()]))] hosted_keys: Option>, - #[values("os.environ/KEY", "os.environ/oidc/env/KEY")] name: &str, - ) { - let server = MockServer::start().await; - let expected = name == "os.environ/KEY" - && access_mode.readable() - && hosted_keys - .as_ref() - .is_none_or(|keys| keys.iter().any(|key| key == "KEY")); - Mock::given(method("POST")) - .respond_with( - ResponseTemplate::new(200) - .set_body_json(serde_json::json!({"SecretString":"remote"})), - ) - .expect(u64::from(expected)) - .mount(&server) - .await; - let state = state( - &server, +#[rstest::rstest] +#[case::lowercase("true", false)] +#[case::capitalized("True", true)] +#[case::literal("(True)", true)] +#[tokio::test] +async fn excluded_hosted_keys_keep_the_python_manager_conversion_path( + #[case] raw: &'static str, + #[case] boolean: bool, +) { + let resolver = SecretResolver::new_python_compatible( + Arc::new(SecretManagerState::new( + SecretManager::External(Arc::new(FixedManager::custom(Err(())))), KeyManagementSettings { - access_mode, - hosted_keys, + hosted_keys: Some(vec!["OTHER".into()]), ..Default::default() }, - ); - assert!(state.backend().is_some()); - assert_eq!(state.settings().unwrap().access_mode, access_mode); - assert_eq!(secret_manager_would_be_consulted(&state, name), expected); - let resolver = SecretResolver::new( - Arc::new(state), - Arc::new(|_: &str| Some("environment".into())), - OidcResolver::default(), - ); - assert_eq!( - resolver - .get_secret_str(name, None) - .await - .unwrap() - .unwrap() - .expose(), - if expected { "remote" } else { "environment" } - ); - } + )), + Arc::new(move |_: &str| Some(raw.to_owned())), + OidcResolver::default(), + ); + assert_eq!( + resolver.get_secret("KEY", None).await.unwrap(), + Some(if boolean { + Secret::Bool(true) + } else { + Secret::String(SecretValue::new(raw)) + }) + ); } -#[cfg(feature = "google")] #[rstest::rstest] -#[case::missing(404)] -#[case::failure(503)] +#[case::missing(Ok(None), None)] +#[case::empty( + Ok(Some(Secret::String(SecretValue::new("")))), + Some(Secret::String(SecretValue::new(""))) +)] +#[case::failed(Err(()), Some(Secret::String(SecretValue::new("environment"))))] #[tokio::test] -async fn google_resolver_distinguishes_absence_from_failure(#[case] status: u16) { - use litellm_secrets::{ - FailurePolicy, KeyManagementSettings, SecretManager, google::GoogleSecretManager, - }; - use wiremock::{Mock, MockServer, ResponseTemplate, matchers::method}; - let server = MockServer::start().await; - Mock::given(method("GET")) - .respond_with(ResponseTemplate::new(status)) - .expect(2) - .mount(&server) - .await; - let environment: Arc = - Arc::new(|name: &str| match name { - "VERTEX_AI_API_KEY" => Some("token".into()), - "KEY" => Some("environment".into()), - _ => None, - }); - let manager = GoogleSecretManager::with_client( - reqwest::Client::new(), - server.uri().parse().unwrap(), - "project".into(), - environment.clone(), - None, - false, - ) - .unwrap(); - let state = SecretManagerState::new( - SecretManager::GoogleSecretManager(manager), - KeyManagementSettings::default(), +async fn azure_callback_absence_preserves_none_but_errors_fall_back( + #[case] reply: Result, ()>, + #[case] expected: Option, +) { + let resolver = SecretResolver::new_python_compatible( + Arc::new(SecretManagerState::new( + SecretManager::External(Arc::new(FixedManager { + reply, + system: KeyManagementSystem::AzureKeyVault, + })), + KeyManagementSettings::default(), + )), + Arc::new(|_: &str| Some("environment".into())), + OidcResolver::default(), ); - let resolver = SecretResolver::new(Arc::new(state), environment, OidcResolver::default()); - let result = resolver.get_secret_str("KEY", None).await; - if status == 404 { - assert_eq!(result.unwrap().unwrap().expose(), "environment"); - } else { - assert!( - matches!(result, Err(Error::Google(litellm_secrets::google::Error::Status(actual))) if actual == status) - ); - } assert_eq!( resolver - .with_failure_policy(FailurePolicy::EnvironmentFallback) - .get_secret_str("KEY", None) + .get_secret("key", Some(Secret::String(SecretValue::new("default")))) .await - .unwrap() - .unwrap() - .expose(), - "environment" + .unwrap(), + expected ); } diff --git a/litellm-rust/crates/secrets/tests/source.rs b/litellm-rust/crates/secrets/tests/source.rs new file mode 100644 index 00000000000..b4782c6af86 --- /dev/null +++ b/litellm-rust/crates/secrets/tests/source.rs @@ -0,0 +1,62 @@ +#[cfg(test)] +mod tests { + use rstest::rstest; + + use litellm_secrets::source::{EnvironmentSecrets, SecretSource}; + + #[rstest] + #[case::lowercase_true("LITELLM_ENVIRONMENT_SECRETS_TRUE", "true", None)] + #[case::padded_false("LITELLM_ENVIRONMENT_SECRETS_FALSE", " FALSE ", None)] + #[case::text("LITELLM_ENVIRONMENT_SECRETS_TEXT", "secret", Some("secret"))] + #[tokio::test] + async fn python_environment_values_are_absent_like_get_secret_str( + #[case] name: &'static str, + #[case] value: &str, + #[case] expected: Option<&str>, + ) { + unsafe { std::env::set_var(name, value) }; + let secret = EnvironmentSecrets::python_compatible() + .resolve(&[name]) + .await + .unwrap() + .get(name); + unsafe { std::env::remove_var(name) }; + assert_eq!(secret.as_deref(), expected); + } +} + +#[tokio::test] +async fn dynamic_names_use_the_same_resolver_and_snapshots_never_do_fresh_lookups() { + use litellm_secrets::source::SecretSource; + use litellm_secrets::{OidcResolver, SecretManagerState, SecretResolver}; + use std::sync::{ + Arc, + atomic::{AtomicUsize, Ordering}, + }; + + let calls = Arc::new(AtomicUsize::new(0)); + let reads = calls.clone(); + let source = SecretResolver::new( + Arc::new(SecretManagerState::default()), + Arc::new(move |name: &str| { + reads.fetch_add(1, Ordering::SeqCst); + (name != "missing").then(|| name.to_owned()) + }), + OidcResolver::default(), + ); + let snapshot = source.resolve(&["declared", "missing"]).await.unwrap(); + let name = format!("runtime-{}", "key"); + assert_eq!(snapshot.get("declared").as_deref(), Some("declared")); + assert_eq!(snapshot.get("missing"), None); + assert_eq!(snapshot.get(&name), None); + assert_eq!(calls.load(Ordering::SeqCst), 2); + assert_eq!( + SecretSource::get_secret_str(&source, &name) + .await + .unwrap() + .unwrap() + .expose(), + name + ); + assert_eq!(calls.load(Ordering::SeqCst), 3); +} diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 92a75bf953a..9a42ee75c51 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -6344,10 +6344,12 @@ class ProxyConfig: ### [DEPRECATED] LOAD FROM GOOGLE KMS ### old way of loading from google kms use_google_kms: Final = general_settings.get("use_google_kms", False) - load_google_kms(use_google_kms=use_google_kms) + if use_google_kms: + self.initialize_secret_manager(KeyManagementSystem.GOOGLE_KMS.value) ### [DEPRECATED] LOAD FROM AZURE KEY VAULT ### old way of loading from azure secret manager use_azure_key_vault: Final = general_settings.get("use_azure_key_vault", False) - load_from_azure_key_vault(use_azure_key_vault=use_azure_key_vault) + if use_azure_key_vault is not False: + self.initialize_secret_manager(KeyManagementSystem.AZURE_KEY_VAULT.value) ### ALERTING ### self._load_alerting_settings(general_settings=general_settings) ### PLUGINS ### @@ -6848,6 +6850,7 @@ class ProxyConfig: """ Initialize the relevant secret manager if `key_management_system` is provided """ + previous_client: Final[object] = litellm.secret_manager_client if key_management_system is not None: if key_management_system == KeyManagementSystem.AZURE_KEY_VAULT.value: ### LOAD FROM AZURE KEY VAULT ### @@ -6896,6 +6899,11 @@ class ProxyConfig: else: raise ValueError("Invalid Key Management System selected") + from litellm.rust_bridge.secret_manager import capture_secret_manager + + if litellm.secret_manager_client is not previous_client: + capture_secret_manager(litellm.secret_manager_client, key_management_system) + def get_model_info_with_id(self, model, db_model=False) -> RouterModelInfo: """ Common logic across add + delete router models diff --git a/litellm/rust_bridge/_native.pyi b/litellm/rust_bridge/_native.pyi index 355a5e5062e..60d2e6224c0 100644 --- a/litellm/rust_bridge/_native.pyi +++ b/litellm/rust_bridge/_native.pyi @@ -2,6 +2,9 @@ from asyncio import Future from collections.abc import AsyncIterator, Coroutine, Iterator, Mapping, Sequence from typing import Never, final +import httpx +from pydantic import JsonValue + from litellm.llms.base_llm.ocr.transformation import OCRResponse from litellm.rust_bridge.messages.entrypoints import LiteLLMMessagesRequest from litellm.rust_bridge.ocr.entrypoints import LiteLLMOcrRequest @@ -371,3 +374,42 @@ __all__ = [ "reserve_process_for_forking", "transcription", ] + +@final +class _SecretManagerRuntime: + @staticmethod + def from_config( + system: str, + environment: Mapping[str, str], + settings: Mapping[str, object] | None = None, + enterprise_enabled: bool = False, + ) -> _SecretManagerRuntime: ... + @staticmethod + def from_client(client: object) -> _SecretManagerRuntime | None: ... + @property + def system(self) -> str: ... + def read_secret(self, name: str, settings: Mapping[str, object] | None = None) -> JsonValue: ... + def read_secret_async(self, name: str, settings: Mapping[str, object] | None = None) -> Future[JsonValue]: ... + def async_write_secret( + self, secret_name: str, secret_value: str, description: str | None = None, + optional_params: Mapping[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, tags: object = None, + ) -> Future[dict[str, JsonValue]]: ... + def async_delete_secret( + self, secret_name: str, recovery_window_in_days: int | None = None, + optional_params: Mapping[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> Future[dict[str, JsonValue]]: ... + def async_rotate_secret( + self, current_secret_name: str, new_secret_name: str, new_secret_value: str, + optional_params: Mapping[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> Future[dict[str, JsonValue]]: ... + def sync_read_secret( + self, secret_name: str, optional_params: Mapping[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, primary_secret_name: str | None = None, + ) -> JsonValue: ... + def async_read_secret( + self, secret_name: str, optional_params: Mapping[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, primary_secret_name: str | None = None, + ) -> Future[JsonValue]: ... diff --git a/litellm/rust_bridge/secret_manager.py b/litellm/rust_bridge/secret_manager.py new file mode 100644 index 00000000000..ca7dc2e6434 --- /dev/null +++ b/litellm/rust_bridge/secret_manager.py @@ -0,0 +1,285 @@ +from __future__ import annotations + +import os +from collections.abc import Awaitable, Mapping +from dataclasses import dataclass, field +from importlib import import_module +from typing import Final, Protocol, runtime_checkable + +import httpx +from pydantic import JsonValue + +from litellm.rust_bridge.bindings import NativeBinding +from litellm.rust_bridge.catalog import Rules, SecretManagerContext, decision +from litellm.rust_bridge.configuration import Decision +from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem + + +@dataclass(frozen=True, slots=True) +class NativeSecretManagerConfig: + system: str + environment: tuple[tuple[str, str], ...] = field(repr=False) + settings: Mapping[str, object] = field(repr=False) + enterprise_enabled: bool + owner_type: type[object] + environment_attributes: tuple[tuple[str, str], ...] + settings_attributes: tuple[str, ...] + methods: tuple[tuple[str, object], ...] = field(repr=False) + + +@dataclass(frozen=True, slots=True) +class _ClientAdapter: + system: KeyManagementSystem + module: str + name: str + methods: tuple[str, ...] + environment_attributes: tuple[tuple[str, str], ...] = () + settings_attributes: tuple[str, ...] = () + enterprise_enabled: bool = False + + +_ADAPTERS: Final = ( + _ClientAdapter( + KeyManagementSystem.AWS_SECRET_MANAGER, + "litellm.secret_managers.aws_secret_manager_v2", + "AWSSecretsManagerV2", + ("sync_read_secret", "async_read_secret"), + settings_attributes=( + "aws_region_name", + "aws_role_name", + "aws_session_name", + "aws_external_id", + "aws_profile_name", + "aws_web_identity_token", + "aws_sts_endpoint", + "replica_regions", + "kms_key_id", + ), + ), + _ClientAdapter( + KeyManagementSystem.HASHICORP_VAULT, + "litellm.secret_managers.hashicorp_secret_manager", + "HashicorpSecretManager", + ("sync_read_secret", "async_read_secret", "async_write_secret", "async_delete_secret", "async_rotate_secret"), + environment_attributes=( + ("HCP_VAULT_ADDR", "vault_addr"), + ("HCP_VAULT_TOKEN", "vault_token"), + ("HCP_VAULT_NAMESPACE", "vault_namespace"), + ("HCP_VAULT_LOGIN_NAMESPACE", "login_namespace_override"), + ("HCP_VAULT_SECRET_NAMESPACE", "secret_namespace_override"), + ("HCP_VAULT_MOUNT_NAME", "vault_mount_name"), + ("HCP_VAULT_PATH_PREFIX", "vault_path_prefix"), + ("HCP_VAULT_CLIENT_CERT", "tls_cert_path"), + ("HCP_VAULT_CLIENT_KEY", "tls_key_path"), + ("HCP_VAULT_CERT_ROLE", "vault_cert_role"), + ("HCP_VAULT_APPROLE_ROLE_ID", "approle_role_id"), + ("HCP_VAULT_APPROLE_SECRET_ID", "approle_secret_id"), + ("HCP_VAULT_APPROLE_MOUNT_PATH", "approle_mount_path"), + ("HCP_VAULT_REFRESH_INTERVAL", "cache.default_ttl"), + ), + enterprise_enabled=True, + ), + _ClientAdapter( + KeyManagementSystem.CYBERARK, + "litellm.secret_managers.cyberark_secret_manager", + "CyberArkSecretManager", + ("sync_read_secret", "async_read_secret", "async_write_secret", "async_delete_secret", "async_rotate_secret"), + environment_attributes=( + ("CYBERARK_API_BASE", "conjur_addr"), + ("CYBERARK_ACCOUNT", "conjur_account"), + ("CYBERARK_USERNAME", "conjur_username"), + ("CYBERARK_API_KEY", "conjur_api_key"), + ("CYBERARK_CLIENT_CERT", "tls_cert_path"), + ("CYBERARK_CLIENT_KEY", "tls_key_path"), + ("CYBERARK_SSL_VERIFY", "ssl_verify"), + ("CYBERARK_REFRESH_INTERVAL", "cache.default_ttl"), + ), + enterprise_enabled=True, + ), + _ClientAdapter( + KeyManagementSystem.GOOGLE_SECRET_MANAGER, + "litellm.secret_managers.google_secret_manager", + "GoogleSecretManager", + ("get_secret_from_google_secret_manager",), + environment_attributes=( + ("GOOGLE_SECRET_MANAGER_PROJECT_ID", "PROJECT_ID"), + ("GOOGLE_SECRET_MANAGER_REFRESH_INTERVAL", "cache.default_ttl"), + ("GOOGLE_SECRET_MANAGER_ALWAYS_READ_SECRET_MANAGER", "always_read_secret_manager"), + ), + enterprise_enabled=True, + ), +) + +_SDK_ADAPTERS: Final = ( + _ClientAdapter(KeyManagementSystem.AZURE_KEY_VAULT, "azure.keyvault.secrets", "SecretClient", ("get_secret",)), + _ClientAdapter(KeyManagementSystem.GOOGLE_KMS, "google.cloud.kms_v1", "KeyManagementServiceClient", ("decrypt",)), +) + + +def _capture(client: object, adapter: _ClientAdapter) -> NativeSecretManagerConfig: + prefixes: Final = ("AWS_", "AZURE_", "GOOGLE_", "VERTEX_", "GCS_", "HCP_VAULT_", "CYBERARK_") + config: Final = NativeSecretManagerConfig( + system=adapter.system.value, + environment=tuple( + (name, value) + for name, value in os.environ.items() + if name.startswith(prefixes) or name == "SECRET_MANAGER_REFRESH_INTERVAL" + ), + settings=KeyManagementSettings().model_dump(mode="json"), + enterprise_enabled=adapter.enterprise_enabled, + owner_type=type(client), + environment_attributes=adapter.environment_attributes, + settings_attributes=adapter.settings_attributes, + methods=tuple((name, getattr(type(client), name)) for name in adapter.methods), + ) + vars(client)["_litellm_native_secret_config"] = config + return config + + +def capture_secret_manager(client: object, system: str) -> None: + for adapter in (*_ADAPTERS, *_SDK_ADAPTERS): + if ( + adapter.system.value == system + and type(client).__name__ == adapter.name + and type(client) is getattr(import_module(adapter.module), adapter.name) + ): + _capture(client, adapter) + return + if ( + system == KeyManagementSystem.AWS_KMS.value + and type(client).__module__ == "botocore.client" + and type(client).__name__ == "KMS" + ): + _capture(client, _ClientAdapter(KeyManagementSystem.AWS_KMS, "botocore.client", "KMS", ("decrypt",))) + + +def native_secret_manager_config(client: object) -> NativeSecretManagerConfig | None: + captured: Final = getattr(client, "_litellm_native_secret_config", None) + if isinstance(captured, NativeSecretManagerConfig): + return captured + for adapter in _ADAPTERS: + if type(client).__module__ == adapter.module and type(client) is getattr( + import_module(adapter.module), adapter.name + ): + return _capture(client, adapter) + return None + + +class NativeSecretManagerRuntime(Protocol): + @property + def system(self) -> str: ... + + def read_secret(self, name: str, settings: Mapping[str, object] | None = None) -> JsonValue: ... + + +@runtime_checkable +class NativeSecretManagerFactory(Protocol): + @staticmethod + def from_client(client: object) -> NativeSecretManagerRuntime | None: ... + + +def _factory(value: object) -> NativeSecretManagerFactory | None: + return value if isinstance(value, NativeSecretManagerFactory) and callable(value.from_client) else None + + +NATIVE_SECRET_MANAGER: Final = NativeBinding("_SecretManagerRuntime", validate=_factory) + + +def resolve_native_secret_manager( + client: object, + system: str, + rules: Rules | None = None, + *, + binding: NativeBinding[NativeSecretManagerFactory] = NATIVE_SECRET_MANAGER, +) -> NativeSecretManagerRuntime | None: + if system in ("custom", "local"): + return None + selected: Final = decision(SecretManagerContext(system=system), rules) + if selected is Decision.PYTHON: + return None + factory: Final = binding.load() + if factory is None: + if selected is Decision.RUST_REQUIRED: + raise RuntimeError("Rust secret manager runtime is unavailable") + return None + runtime: Final = factory.from_client(client) + if runtime is not None and runtime.system != system: + raise ValueError("Native secret manager system does not match configuration") + return runtime + + +@runtime_checkable +class NativeProviderReader(Protocol): + def sync_read_secret( + self, + secret_name: str, + optional_params: Mapping[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> str | None: ... + + def async_read_secret( + self, + secret_name: str, + optional_params: Mapping[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> Awaitable[str | None]: ... + + +def resolve_native_provider_reader( + client: object, + system: str, + rules: Rules | None = None, + *, + binding: NativeBinding[NativeSecretManagerFactory] = NATIVE_SECRET_MANAGER, +) -> NativeProviderReader | None: + runtime: Final = resolve_native_secret_manager(client, system, rules, binding=binding) + if runtime is None: + return None + if not isinstance(runtime, NativeProviderReader): + raise TypeError("Rust secret manager provider reads are unavailable") + return runtime + + +@runtime_checkable +class NativeProviderWriter(Protocol): + def async_write_secret( + self, + secret_name: str, + secret_value: str, + description: str | None = None, + optional_params: Mapping[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + tags: object = None, + ) -> Awaitable[dict[str, JsonValue]]: ... + + def async_delete_secret( + self, + secret_name: str, + recovery_window_in_days: int | None = 7, + optional_params: Mapping[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> Awaitable[dict[str, JsonValue]]: ... + + def async_rotate_secret( + self, + current_secret_name: str, + new_secret_name: str, + new_secret_value: str, + optional_params: Mapping[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> Awaitable[dict[str, JsonValue]]: ... + + +def resolve_native_provider_writer( + client: object, + system: str, + rules: Rules | None = None, + *, + binding: NativeBinding[NativeSecretManagerFactory] = NATIVE_SECRET_MANAGER, +) -> NativeProviderWriter | None: + runtime: Final = resolve_native_secret_manager(client, system, rules, binding=binding) + if runtime is None: + return None + if not isinstance(runtime, NativeProviderWriter): + raise TypeError("Rust secret manager provider writes are unavailable") + return runtime diff --git a/litellm/rust_bridge/settings.py b/litellm/rust_bridge/settings.py index 46e96b60620..7861f50574f 100644 --- a/litellm/rust_bridge/settings.py +++ b/litellm/rust_bridge/settings.py @@ -1,7 +1,10 @@ from __future__ import annotations from dataclasses import dataclass -from typing import Final +from typing import TYPE_CHECKING, Final + +if TYPE_CHECKING: + from litellm.rust_bridge.catalog import Rules @dataclass(frozen=True, slots=True) @@ -34,6 +37,7 @@ class ProviderDefaults: @dataclass(frozen=True, slots=True) class SecretManager: readable: bool + native: bool @dataclass(frozen=True, slots=True) @@ -58,12 +62,24 @@ class SecretManagerBinding: settings_object: object -def secret_manager() -> SecretManager: +def secret_manager(rules: Rules | None = None) -> SecretManager: + import litellm + from litellm.rust_bridge.catalog import SecretManagerContext, decision + from litellm.rust_bridge.configuration import Decision from litellm.secret_managers.main import ( _should_read_secret_from_secret_manager, # pyright: ignore[reportPrivateUsage] # canonical resolver is private ) - return SecretManager(readable=_should_read_secret_from_secret_manager()) + readable: Final = _should_read_secret_from_secret_manager() + system: Final = ( + litellm._key_management_system # pyright: ignore[reportPrivateUsage] # canonical key management globals are private + ) + native: Final = ( + readable + and system is not None + and decision(SecretManagerContext(system=system.value), rules) is not Decision.PYTHON + ) + return SecretManager(readable=readable, native=native) def secret_manager_binding() -> SecretManagerBinding: diff --git a/litellm/secret_managers/aws_secret_manager_v2.py b/litellm/secret_managers/aws_secret_manager_v2.py index 3e9f4e259d5..80fe3f38c03 100644 --- a/litellm/secret_managers/aws_secret_manager_v2.py +++ b/litellm/secret_managers/aws_secret_manager_v2.py @@ -31,6 +31,7 @@ from litellm.llms.custom_httpx.http_handler import ( get_async_httpx_client, ) from litellm.proxy._types import KeyManagementSystem +from litellm.rust_bridge.secret_manager import resolve_native_provider_reader from litellm.secret_managers.main import get_secret_str from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.secret_managers.main import KeyManagementSettings @@ -140,6 +141,10 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): secret_name=secret_name, primary_secret_name=primary_secret_name ) + native: Final = resolve_native_provider_reader(self, "aws_secret_manager") + if native is not None: + return await native.async_read_secret(secret_name, optional_params, timeout) + endpoint_url, headers, body = self._prepare_request( action="GetSecretValue", secret_name=secret_name, @@ -192,6 +197,10 @@ class AWSSecretsManagerV2(BaseAWSLLM, BaseSecretManager): secret_name=secret_name, primary_secret_name=primary_secret_name ) + native: Final = resolve_native_provider_reader(self, "aws_secret_manager") + if native is not None: + return native.sync_read_secret(secret_name, optional_params, timeout) + endpoint_url, headers, body = self._prepare_request( action="GetSecretValue", secret_name=secret_name, diff --git a/litellm/secret_managers/cyberark_secret_manager.py b/litellm/secret_managers/cyberark_secret_manager.py index a49349991c5..b28e15c4446 100644 --- a/litellm/secret_managers/cyberark_secret_manager.py +++ b/litellm/secret_managers/cyberark_secret_manager.py @@ -15,6 +15,7 @@ from litellm.llms.custom_httpx.http_handler import ( httpxSpecialProvider, ) from litellm.proxy._types import KeyManagementSystem +from litellm.rust_bridge.secret_manager import resolve_native_provider_reader, resolve_native_provider_writer from .base_secret_manager import BaseSecretManager, raise_if_unsafe_secret_name from .main import str_to_bool @@ -186,6 +187,10 @@ class CyberArkSecretManager(BaseSecretManager): Returns: Optional[str]: The secret value if found, None otherwise """ + native: Final = resolve_native_provider_reader(self, "cyberark") + if native is not None: + return await native.async_read_secret(secret_name, optional_params, timeout) + # Check cache first if self.cache.get_cache(secret_name) is not None: return self.cache.get_cache(secret_name) @@ -232,6 +237,10 @@ class CyberArkSecretManager(BaseSecretManager): Returns: Optional[str]: The secret value if found, None otherwise """ + native: Final = resolve_native_provider_reader(self, "cyberark") + if native is not None: + return native.sync_read_secret(secret_name, optional_params, timeout) + # Check cache first if self.cache.get_cache(secret_name) is not None: return self.cache.get_cache(secret_name) @@ -281,6 +290,12 @@ class CyberArkSecretManager(BaseSecretManager): Returns: dict: Response containing status and details of the operation """ + native: Final = resolve_native_provider_writer(self, "cyberark") + if native is not None: + return await native.async_write_secret( + secret_name, secret_value, description, optional_params, timeout, tags + ) + async_client: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.SecretManager, params={"ssl_verify": self.ssl_verify}, @@ -326,6 +341,10 @@ class CyberArkSecretManager(BaseSecretManager): Returns: dict: Response indicating operation not supported """ + native: Final = resolve_native_provider_writer(self, "cyberark") + if native is not None: + return await native.async_delete_secret(secret_name, recovery_window_in_days, optional_params, timeout) + verbose_logger.warning( "CyberArk Conjur does not support direct secret deletion. Secrets must be removed through policy updates." ) @@ -337,3 +356,28 @@ class CyberArkSecretManager(BaseSecretManager): "status": "not_supported", "message": "CyberArk Conjur does not support direct secret deletion. Use policy updates to remove variables.", } + + async def async_rotate_secret( + self, + current_secret_name: str, + new_secret_name: str, + new_secret_value: str, + optional_params: dict | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> dict: + native: Final = resolve_native_provider_writer(self, "cyberark") + if native is not None: + return await native.async_rotate_secret( + current_secret_name, + new_secret_name, + new_secret_value, + optional_params, + timeout, + ) + return await super().async_rotate_secret( + current_secret_name, + new_secret_name, + new_secret_value, + optional_params, + timeout, + ) diff --git a/litellm/secret_managers/dispatch.py b/litellm/secret_managers/dispatch.py new file mode 100644 index 00000000000..4912771eb3c --- /dev/null +++ b/litellm/secret_managers/dispatch.py @@ -0,0 +1,30 @@ +from typing import Final + +from pydantic import JsonValue + +from litellm.rust_bridge.bindings import NativeBinding +from litellm.rust_bridge.catalog import Rules +from litellm.rust_bridge.secret_manager import ( + NATIVE_SECRET_MANAGER, + NativeSecretManagerFactory, + resolve_native_secret_manager, +) +from litellm.secret_managers.secret_manager_handler import get_secret_from_manager as python_get_secret_from_manager +from litellm.types.secret_managers.main import KeyManagementSettings + + +def get_secret_from_manager( + client: object, + key_manager: str, + secret_name: str, + key_management_settings: KeyManagementSettings | None = None, + *, + rules: Rules | None = None, + binding: NativeBinding[NativeSecretManagerFactory] = NATIVE_SECRET_MANAGER, +) -> JsonValue: + native: Final = resolve_native_secret_manager(client, key_manager, rules, binding=binding) + if native is None: + return python_get_secret_from_manager(client, key_manager, secret_name, key_management_settings) + return native.read_secret( + secret_name, key_management_settings.model_dump(mode="json") if key_management_settings is not None else None + ) diff --git a/litellm/secret_managers/google_secret_manager.py b/litellm/secret_managers/google_secret_manager.py index 6674549e39c..913fd4d5204 100644 --- a/litellm/secret_managers/google_secret_manager.py +++ b/litellm/secret_managers/google_secret_manager.py @@ -9,6 +9,7 @@ from litellm.constants import SECRET_MANAGER_REFRESH_INTERVAL from litellm.integrations.gcs_bucket.gcs_bucket_base import GCSBucketBase from litellm.llms.custom_httpx.http_handler import _get_httpx_client from litellm.proxy._types import CommonProxyErrors, KeyManagementSystem +from litellm.rust_bridge.secret_manager import resolve_native_provider_reader class GoogleSecretManager(GCSBucketBase): @@ -60,6 +61,10 @@ class GoogleSecretManager(GCSBucketBase): Returns: str: The secret value if successful, None otherwise. """ + native: Final = resolve_native_provider_reader(self, "google_secret_manager") + if native is not None: + return native.sync_read_secret(secret_name) + if self.always_read_secret_manager is not True: cached_secret: Final = self.cache.get_cache(secret_name) if cached_secret is not None: diff --git a/litellm/secret_managers/hashicorp_secret_manager.py b/litellm/secret_managers/hashicorp_secret_manager.py index e37a912c7e1..27523892d61 100644 --- a/litellm/secret_managers/hashicorp_secret_manager.py +++ b/litellm/secret_managers/hashicorp_secret_manager.py @@ -16,6 +16,7 @@ from litellm.llms.custom_httpx.http_handler import ( httpxSpecialProvider, ) from litellm.proxy._types import KeyManagementSystem +from litellm.rust_bridge.secret_manager import resolve_native_provider_reader, resolve_native_provider_writer from .base_secret_manager import BaseSecretManager, raise_if_unsafe_secret_name @@ -405,6 +406,10 @@ class HashicorpSecretManager(BaseSecretManager): secret_name is just the path inside the KV mount (e.g., 'myapp/config'). Returns the entire data dict from data.data, or None on failure. """ + native: Final = resolve_native_provider_reader(self, "hashicorp_vault") + if native is not None: + return await native.async_read_secret(secret_name, optional_params, timeout) + async_client: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.SecretManager, ) @@ -436,6 +441,10 @@ class HashicorpSecretManager(BaseSecretManager): secret_name is just the path inside the KV mount (e.g., 'myapp/config'). Returns the entire data dict from data.data, or None on failure. """ + native: Final = resolve_native_provider_reader(self, "hashicorp_vault") + if native is not None: + return native.sync_read_secret(secret_name, optional_params, timeout) + sync_client: Final = _get_httpx_client() try: target: Final = self._build_secret_target(secret_name, optional_params) @@ -476,6 +485,12 @@ class HashicorpSecretManager(BaseSecretManager): Returns: dict: Response containing status and details of the operation """ + native: Final = resolve_native_provider_writer(self, "hashicorp_vault") + if native is not None: + return await native.async_write_secret( + secret_name, secret_value, description, optional_params, timeout, tags + ) + async_client: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.SecretManager, params={"timeout": timeout}, @@ -525,6 +540,12 @@ class HashicorpSecretManager(BaseSecretManager): On success, returns the response from async_write_secret. On error, returns {"status": "error", "message": "error message"} """ + native: Final = resolve_native_provider_writer(self, "hashicorp_vault") + if native is not None: + return await native.async_rotate_secret( + current_secret_name, new_secret_name, new_secret_value, optional_params, timeout + ) + async_client: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.SecretManager, params={"timeout": timeout}, @@ -671,6 +692,10 @@ class HashicorpSecretManager(BaseSecretManager): Returns: dict: Response containing status and details of the operation """ + native: Final = resolve_native_provider_writer(self, "hashicorp_vault") + if native is not None: + return await native.async_delete_secret(secret_name, recovery_window_in_days, optional_params, timeout) + async_client: Final = get_async_httpx_client( llm_provider=httpxSpecialProvider.SecretManager, params={"timeout": timeout}, diff --git a/litellm/secret_managers/main.py b/litellm/secret_managers/main.py index e89fbbdab65..f09ddd1d5a9 100644 --- a/litellm/secret_managers/main.py +++ b/litellm/secret_managers/main.py @@ -12,10 +12,10 @@ import litellm from litellm._logging import verbose_logger from litellm.caching.caching import DualCache from litellm.llms.custom_httpx.http_handler import HTTPHandler +from litellm.secret_managers.dispatch import get_secret_from_manager from litellm.secret_managers.get_azure_ad_token_provider import ( get_azure_ad_token_provider, ) -from litellm.secret_managers.secret_manager_handler import get_secret_from_manager oidc_cache: Final = DualCache() diff --git a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py index 89cb8356289..7e198bc9131 100644 --- a/tests/test_litellm/proxy/proxy_server/test_proxy_config.py +++ b/tests/test_litellm/proxy/proxy_server/test_proxy_config.py @@ -16,6 +16,7 @@ import re from collections.abc import Mapping from dataclasses import dataclass from datetime import datetime +from pathlib import Path from types import MappingProxyType, SimpleNamespace from typing import Any, Dict, Final from unittest.mock import AsyncMock, MagicMock @@ -2072,6 +2073,88 @@ def test_ProxyConfig__load_environment_variables_blocks_dangerous_keys(monkeypat # --------------------------------------------------------------------------- +@pytest.mark.asyncio +@pytest.mark.parametrize( + ("flag", "system"), (("use_google_kms", "google_kms"), ("use_azure_key_vault", "azure_key_vault")) +) +async def test_load_config_legacy_secret_manager_flags_capture_the_initialized_client( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, flag: str, system: str +) -> None: + if system == "azure_key_vault": + client_type: Final = pytest.importorskip("azure.keyvault.secrets").SecretClient + else: + client_type: Final = pytest.importorskip("google.cloud.kms_v1").KeyManagementServiceClient + + from litellm.rust_bridge.secret_manager import native_secret_manager_config + + credentials_file: Final = tmp_path / "credentials.json" + credentials_file.write_text( + json.dumps( + { + "type": "authorized_user", + "client_id": "test-client", + "client_secret": "test-secret", + "refresh_token": "test", + } + ) + ) + config_file: Final = tmp_path / "legacy-secret-manager.yaml" + config_file.write_text( + f"model_list: []\ngeneral_settings:\n {flag}: true\n key_management_settings:\n access_mode: write_only\n" + ) + monkeypatch.setenv("GOOGLE_APPLICATION_CREDENTIALS", str(credentials_file)) + monkeypatch.setenv("GOOGLE_KMS_RESOURCE_NAME", "projects/test/locations/global/keyRings/test/cryptoKeys/test") + monkeypatch.setenv("AZURE_KEY_VAULT_URI", "https://test.vault.azure.net") + monkeypatch.setattr(litellm, "secret_manager_client", None) + monkeypatch.setattr(litellm, "_key_management_system", None) + monkeypatch.setattr(litellm, "_google_kms_resource_name", None) + monkeypatch.setattr(litellm, "_key_management_settings", litellm._key_management_settings) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) + + await ProxyConfig().load_config(router=None, config_file_path=str(config_file)) + + client: Final = litellm.secret_manager_client + assert isinstance(client, client_type) + try: + captured: Final = native_secret_manager_config(client) + assert captured is not None + assert captured.system == system + assert dict(captured.environment)["GOOGLE_APPLICATION_CREDENTIALS"] == str(credentials_file) + assert litellm._key_management_system is not None + assert litellm._key_management_system.value == system + finally: + if system == "azure_key_vault": + client.close() + else: + client.transport.close() + + +@pytest.mark.asyncio +@pytest.mark.parametrize("flag", ("null", "false")) +async def test_load_config_disabled_google_kms_does_not_initialize_a_manager( + tmp_path: Path, monkeypatch: pytest.MonkeyPatch, flag: str +) -> None: + config_file: Final = tmp_path / "disabled-kms.yaml" + config_file.write_text(f"model_list: []\ngeneral_settings:\n use_google_kms: {flag}\n") + monkeypatch.setattr(litellm, "secret_manager_client", None) + monkeypatch.setattr(litellm, "_key_management_system", None) + monkeypatch.setattr("litellm.proxy.proxy_server.prisma_client", None) + monkeypatch.setattr("litellm.proxy.proxy_server.store_model_in_db", False) + monkeypatch.delenv("LITELLM_CONFIG_BUCKET_NAME", raising=False) + monkeypatch.delenv("GOOGLE_APPLICATION_CREDENTIALS", raising=False) + + _router, model_list, general_settings = await ProxyConfig().load_config( + router=None, config_file_path=str(config_file) + ) + + assert model_list == [] + assert general_settings["use_google_kms"] is (None if flag == "null" else False) + assert litellm.secret_manager_client is None + assert litellm._key_management_system is None + + @pytest.mark.asyncio async def test_ProxyConfig_load_config_minimal_yaml(tmp_path, monkeypatch): f = tmp_path / "c.yaml" diff --git a/tests/test_litellm/rust_bridge/AGENTS.md b/tests/test_litellm/rust_bridge/AGENTS.md new file mode 100644 index 00000000000..351bd582ce7 --- /dev/null +++ b/tests/test_litellm/rust_bridge/AGENTS.md @@ -0,0 +1,7 @@ +# Rust bridge tests + +Test what each side of the bridge does, not the rollout policy that picks a side. `LITELLM_RUST` and `catalog.RULES` change every time a route or backend rolls forward, so a test that sets the env var or patches the catalog to reach a path goes red on a policy change even when the code under test is fine + +Call each path directly with an explicit decision instead. The Python path is the implementation the dispatcher falls back to, e.g. `litellm.ocr.main.ocr`. The Rust path is the native binding, e.g. `NATIVE_OCR.load()` from `litellm/rust_bridge/ocr/entrypoints.py`, called with the request, args and kwargs that dispatch would hand it. When the native side reads a policy-derived setting such as `settings.secret_manager().native`, pin that field in the test instead of deriving it from the catalog. `ocr/test_secrets.py` shows the pattern + +Rollout policy itself, meaning which rule matches and what `LITELLM_RUST` changes, belongs in `test_catalog.py`, `test_configuration.py` and `test_dispatch.py`, tested against rules the test builds rather than the shipped `catalog.RULES` diff --git a/tests/test_litellm/rust_bridge/ocr/test_secrets.py b/tests/test_litellm/rust_bridge/ocr/test_secrets.py index 085a42dd373..a91c0ff5bc8 100644 --- a/tests/test_litellm/rust_bridge/ocr/test_secrets.py +++ b/tests/test_litellm/rust_bridge/ocr/test_secrets.py @@ -1,6 +1,11 @@ from __future__ import annotations -from typing import Final +import asyncio +from collections.abc import Awaitable, Generator, Mapping +from contextlib import contextmanager +from dataclasses import replace +from types import MappingProxyType +from typing import Final, Literal, Protocol, TypeAlias, cast import httpx import pytest @@ -8,14 +13,27 @@ import pytest import litellm from litellm.integrations.custom_secret_manager import CustomSecretManager from litellm.llms.base_llm.ocr.transformation import OCRResponse -from litellm.rust_bridge import configuration +from litellm.ocr import main +from litellm.rust_bridge import settings +from litellm.rust_bridge.ocr.entrypoints import NATIVE_AOCR, NATIVE_OCR, LiteLLMOcrRequest from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem -from tests.test_litellm_rust.support.recording_server import ResponseSpec, recording_service +from tests.test_litellm_rust.support.recording_server import RecordingServer, ResponseSpec, recording_service +from tests.test_litellm_rust.support.requests import OCR_DOCUMENT, OCR_MODEL, OCR_RESPONSE + +native: Final = pytest.importorskip("litellm.rust_bridge._native") + +AccessMode: TypeAlias = Literal["read_only", "write_only", "read_and_write"] + + +class Ocr(Protocol): + def __call__(self, api_base: str, /) -> Awaitable[OCRResponse]: ... class _VaultSecrets(CustomSecretManager): - def __init__(self) -> None: + def __init__(self, failure: BaseException | None = None) -> None: super().__init__(secret_manager_name="rust_bridge_ocr_test") + self.failure: Final = failure + self.reads: tuple[tuple[str, Mapping[str, object] | None], ...] = () async def async_read_secret( self, @@ -23,7 +41,7 @@ class _VaultSecrets(CustomSecretManager): optional_params: dict[str, object] | None = None, timeout: float | httpx.Timeout | None = None, ) -> str | None: - return "vault-key" if secret_name == "MISTRAL_API_KEY" else None + raise AssertionError("get_secret reads custom managers synchronously") def sync_read_secret( self, @@ -31,82 +49,393 @@ class _VaultSecrets(CustomSecretManager): optional_params: dict[str, object] | None = None, timeout: float | httpx.Timeout | None = None, ) -> str | None: - return "vault-key" if secret_name == "MISTRAL_API_KEY" else None + self.reads = (*self.reads, (secret_name, optional_params)) + if secret_name != "MISTRAL_API_KEY": + return None + if self.failure is not None: + raise self.failure + return "vault-key" + + def key_reads(self) -> tuple[Mapping[str, object] | None, ...]: + return tuple(params for name, params in self.reads if name == "MISTRAL_API_KEY") -async def _call(asynchronous: bool, api_base: str) -> OCRResponse: - if asynchronous: - return await litellm.aocr( - model="mistral/mistral-ocr-latest", - document={"type": "document_url", "document_url": "https://example.com/document.pdf"}, - api_base=api_base, - ) - return litellm.ocr( - model="mistral/mistral-ocr-latest", - document={"type": "document_url", "document_url": "https://example.com/document.pdf"}, +def _native_request(api_base: str) -> LiteLLMOcrRequest: + return LiteLLMOcrRequest( + model=OCR_MODEL, + document=OCR_DOCUMENT, + api_key=None, api_base=api_base, + timeout=None, + custom_llm_provider=None, + extra_headers=None, + kwargs=MappingProxyType({}), ) -_RESPONSE: Final = { - "pages": [{"index": 0, "markdown": "parsed document", "images": []}], - "model": "mistral-ocr-latest", - "usage_info": {"pages_processed": 1}, -} +def _public_kwargs(api_base: str) -> dict[str, object]: + return {"model": OCR_MODEL, "document": OCR_DOCUMENT, "api_base": api_base} -@pytest.mark.asyncio -@pytest.mark.parametrize("asynchronous", (False, True)) -@pytest.mark.parametrize("rust_enabled", ("0", "1")) -@pytest.mark.parametrize("access_mode", ("read_only", "read_and_write")) -@pytest.mark.parametrize("system", (None, KeyManagementSystem.CUSTOM)) -async def test_readable_secret_managers_keep_python_ocr_fallback( - monkeypatch: pytest.MonkeyPatch, - asynchronous: bool, - rust_enabled: str, - access_mode: str, - system: KeyManagementSystem | None, -) -> None: - pytest.importorskip("litellm.rust_bridge._native") - monkeypatch.setenv("LITELLM_RUST", rust_enabled) - monkeypatch.setenv("MISTRAL_API_KEY", "environment-key") - monkeypatch.setattr(litellm, "secret_manager_client", _VaultSecrets()) - monkeypatch.setattr(litellm, "_key_management_system", system) - monkeypatch.setattr( - litellm, - "_key_management_settings", - KeyManagementSettings(access_mode=access_mode, hosted_keys=["MISTRAL_API_KEY"]), - ) - configuration.reset_rust_configuration() +async def _python_ocr(api_base: str) -> OCRResponse: + response: Final = main.ocr(model=OCR_MODEL, document=OCR_DOCUMENT, api_base=api_base) + assert isinstance(response, OCRResponse) + return response + +async def _python_aocr(api_base: str) -> OCRResponse: + return await main.aocr(model=OCR_MODEL, document=OCR_DOCUMENT, api_base=api_base) + + +async def _rust_ocr(api_base: str) -> OCRResponse: + route: Final = NATIVE_OCR.load() + assert route is not None + return route(_native_request(api_base), (), _public_kwargs(api_base)) + + +async def _rust_aocr(api_base: str) -> OCRResponse: + route: Final = NATIVE_AOCR.load() + assert route is not None + return await route(_native_request(api_base), (), _public_kwargs(api_base)) + + +_RUST_PATHS: Final = (_rust_ocr, _rust_aocr) +_RUST_IDS: Final = ("rust-sync", "rust-async") + + +@pytest.fixture(params=(_python_ocr, _python_aocr, *_RUST_PATHS), ids=("python-sync", "python-async", *_RUST_IDS)) +def ocr(request: pytest.FixtureRequest) -> Ocr: + return cast(Ocr, request.param) + + +@pytest.fixture(params=_RUST_PATHS, ids=_RUST_IDS) +def rust_ocr(request: pytest.FixtureRequest) -> Ocr: + return cast(Ocr, request.param) + + +@contextmanager +def _mistral_service(expected_requests: int = 1) -> Generator[RecordingServer]: with recording_service() as server: - server.default_response = ResponseSpec(body=_RESPONSE) - result: Final = await _call(asynchronous, server.base_url) - - assert result.pages[0].markdown == "parsed document" - assert len(server.requests) == 1 - expected_key: Final = "vault-key" if system is KeyManagementSystem.CUSTOM else "environment-key" - assert server.requests[0].headers["authorization"] == f"Bearer {expected_key}" - assert "x-litellm-rust" not in result._hidden_params.get("additional_headers", {}) + server.default_response = ResponseSpec(body=OCR_RESPONSE) + server.expected_requests = expected_requests + yield server -@pytest.mark.asyncio -@pytest.mark.parametrize("asynchronous", (False, True)) -async def test_no_secret_client_leaves_dormant_binding_settings_unread( - monkeypatch: pytest.MonkeyPatch, asynchronous: bool +def _configure( + monkeypatch: pytest.MonkeyPatch, + *, + manager: _VaultSecrets, + key_management: KeyManagementSettings, + native_secret_manager: bool = True, + environment_key: str | None = "environment-key", +) -> None: + if environment_key is None: + monkeypatch.delenv("MISTRAL_API_KEY", raising=False) + else: + monkeypatch.setenv("MISTRAL_API_KEY", environment_key) + monkeypatch.setattr(litellm, "secret_manager_client", manager) + monkeypatch.setattr(litellm, "_key_management_system", KeyManagementSystem.CUSTOM) + monkeypatch.setattr(litellm, "_key_management_settings", key_management) + configured: Final = settings.secret_manager + monkeypatch.setattr(settings, "secret_manager", lambda: replace(configured(), native=native_secret_manager)) + + +@pytest.mark.parametrize( + ("access_mode", "hosted_keys"), + (("read_only", None), ("read_and_write", None), ("read_only", ["MISTRAL_API_KEY"])), +) +async def test_custom_secret_manager_supplies_the_ocr_key( + monkeypatch: pytest.MonkeyPatch, ocr: Ocr, access_mode: AccessMode, hosted_keys: list[str] | None +) -> None: + manager: Final = _VaultSecrets() + key_management: Final = KeyManagementSettings(access_mode=access_mode, hosted_keys=hosted_keys) + _configure(monkeypatch, manager=manager, key_management=key_management) + + with _mistral_service() as server: + await ocr(server.base_url) + + assert server.requests[0].headers["authorization"] == "Bearer vault-key" + assert manager.key_reads(), "the custom manager was never asked for MISTRAL_API_KEY" + assert all(params == key_management.model_dump() for params in manager.key_reads()), manager.key_reads() + + +@pytest.mark.parametrize(("access_mode", "hosted_keys"), (("read_only", ["OTHER"]), ("write_only", None))) +async def test_custom_secret_manager_is_not_read_when_settings_exclude_the_key( + monkeypatch: pytest.MonkeyPatch, ocr: Ocr, access_mode: AccessMode, hosted_keys: list[str] | None +) -> None: + manager: Final = _VaultSecrets() + _configure( + monkeypatch, + manager=manager, + key_management=KeyManagementSettings(access_mode=access_mode, hosted_keys=hosted_keys), + ) + + with _mistral_service() as server: + await ocr(server.base_url) + + assert server.requests[0].headers["authorization"] == "Bearer environment-key" + assert manager.key_reads() == () + + +async def test_custom_secret_manager_exceptions_fall_back_to_the_environment_key( + monkeypatch: pytest.MonkeyPatch, ocr: Ocr +) -> None: + _configure( + monkeypatch, + manager=_VaultSecrets(ValueError("secret manager failed")), + key_management=KeyManagementSettings(access_mode="read_only"), + ) + + with _mistral_service() as server: + await ocr(server.base_url) + + assert server.requests[0].headers["authorization"] == "Bearer environment-key" + + +async def test_custom_secret_manager_exceptions_without_environment_key_raise_missing_key( + monkeypatch: pytest.MonkeyPatch, ocr: Ocr +) -> None: + _configure( + monkeypatch, + manager=_VaultSecrets(ValueError("secret manager failed")), + key_management=KeyManagementSettings(access_mode="read_only"), + environment_key=None, + ) + + with _mistral_service(expected_requests=0) as server: + with pytest.raises(litellm.APIConnectionError, match="Missing Mistral API Key"): + await ocr(server.base_url) + + +async def test_custom_secret_manager_cancellation_propagates_without_provider_io( + monkeypatch: pytest.MonkeyPatch, ocr: Ocr +) -> None: + failure: Final = asyncio.CancelledError("secret manager cancelled") + _configure( + monkeypatch, manager=_VaultSecrets(failure), key_management=KeyManagementSettings(access_mode="read_only") + ) + + with _mistral_service(expected_requests=0) as server: + with pytest.raises(asyncio.CancelledError) as raised: + await ocr(server.base_url) + + assert raised.value is failure + + +async def test_rust_declines_a_readable_secret_manager_it_cannot_resolve( + monkeypatch: pytest.MonkeyPatch, rust_ocr: Ocr +) -> None: + manager: Final = _VaultSecrets() + _configure( + monkeypatch, + manager=manager, + key_management=KeyManagementSettings(access_mode="read_only"), + native_secret_manager=False, + ) + + with _mistral_service(expected_requests=0) as server: + with pytest.raises(native.RustBridgeDeclined): + await rust_ocr(server.base_url) + + assert manager.key_reads() == () + + +async def test_no_secret_client_leaves_dormant_binding_settings_unread( + monkeypatch: pytest.MonkeyPatch, rust_ocr: Ocr ) -> None: - pytest.importorskip("litellm.rust_bridge._native") - monkeypatch.setenv("LITELLM_RUST", "1") monkeypatch.setenv("MISTRAL_API_KEY", "environment-key") monkeypatch.setattr(litellm, "secret_manager_client", None) monkeypatch.setattr(litellm, "_key_management_settings", object()) - configuration.reset_rust_configuration() - with recording_service() as server: - server.default_response = ResponseSpec(body=_RESPONSE) - result: Final = await _call(asynchronous, server.base_url) + with _mistral_service() as server: + await rust_ocr(server.base_url) - assert result.pages[0].markdown == "parsed document" - assert len(server.requests) == 1 assert server.requests[0].headers["authorization"] == "Bearer environment-key" - assert result._hidden_params["additional_headers"]["x-litellm-rust"] == "true" + + +class _FixedSecrets(CustomSecretManager): + def __init__(self, value: str) -> None: + super().__init__(secret_manager_name="rust_bridge_ocr_fixed") + self.value: Final = value + + async def async_read_secret( + self, + secret_name: str, + optional_params: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> str | None: + raise AssertionError("get_secret reads custom managers synchronously") + + def sync_read_secret( + self, + secret_name: str, + optional_params: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> str | None: + return self.value if secret_name == "MISTRAL_API_KEY" else None + + +class _PlainSecretReader: + def sync_read_secret( + self, + secret_name: str, + optional_params: dict[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> str | None: + return "vault-key" + + +class _AzureSecret: + def __init__(self, value: str | None) -> None: + self.value: Final = value + + +def _azure_sdk_client(value: str | None) -> object: + class SecretClient: + def get_secret(self, name: str) -> _AzureSecret: + return _AzureSecret(value if name == "MISTRAL_API_KEY" else None) + + SecretClient.__module__ = "azure.keyvault.secrets._client" + return SecretClient() + + +def _configure_client( + monkeypatch: pytest.MonkeyPatch, + *, + client: object, + system: KeyManagementSystem, + key_management: KeyManagementSettings, + environment_key: str = "environment-key", +) -> None: + monkeypatch.setenv("MISTRAL_API_KEY", environment_key) + monkeypatch.setattr(litellm, "secret_manager_client", client) + monkeypatch.setattr(litellm, "_key_management_system", system) + monkeypatch.setattr(litellm, "_key_management_settings", key_management) + configured: Final = settings.secret_manager + monkeypatch.setattr(settings, "secret_manager", lambda: replace(configured(), native=True)) + + +async def _assert_missing_key(ocr: Ocr) -> None: + with _mistral_service(expected_requests=0) as server: + with pytest.raises(litellm.APIConnectionError, match="Missing Mistral API Key"): + await ocr(server.base_url) + + +@pytest.mark.parametrize("environment_key", ("true", " FALSE ", "True")) +async def test_boolean_environment_keys_count_as_missing( + monkeypatch: pytest.MonkeyPatch, ocr: Ocr, environment_key: str +) -> None: + monkeypatch.setenv("MISTRAL_API_KEY", environment_key) + monkeypatch.setattr(litellm, "secret_manager_client", None) + + await _assert_missing_key(ocr) + + +@pytest.mark.parametrize("manager_key", ("True", "(False)")) +async def test_boolean_manager_keys_count_as_missing( + monkeypatch: pytest.MonkeyPatch, ocr: Ocr, manager_key: str +) -> None: + _configure_client( + monkeypatch, + client=_FixedSecrets(manager_key), + system=KeyManagementSystem.CUSTOM, + key_management=KeyManagementSettings(access_mode="read_only"), + ) + + await _assert_missing_key(ocr) + + +async def test_boolean_environment_fallback_after_a_manager_exception_counts_as_missing( + monkeypatch: pytest.MonkeyPatch, ocr: Ocr +) -> None: + _configure( + monkeypatch, + manager=_VaultSecrets(ValueError("secret manager failed")), + key_management=KeyManagementSettings(access_mode="read_only"), + environment_key="True", + ) + + await _assert_missing_key(ocr) + + +async def test_manager_without_the_key_does_not_fall_back_to_the_environment( + monkeypatch: pytest.MonkeyPatch, ocr: Ocr +) -> None: + _configure_client( + monkeypatch, + client=_azure_sdk_client(None), + system=KeyManagementSystem.AZURE_KEY_VAULT, + key_management=KeyManagementSettings(access_mode="read_only"), + ) + + await _assert_missing_key(ocr) + + +async def test_rust_hosted_keys_exclude_azure_sdk_clients_too(monkeypatch: pytest.MonkeyPatch, rust_ocr: Ocr) -> None: + _configure_client( + monkeypatch, + client=_azure_sdk_client("vault-key"), + system=KeyManagementSystem.AZURE_KEY_VAULT, + key_management=KeyManagementSettings(access_mode="read_only", hosted_keys=["OTHER"]), + ) + + with _mistral_service() as server: + await rust_ocr(server.base_url) + + assert server.requests[0].headers["authorization"] == "Bearer environment-key", ( + "recorded divergence: Python's get_secret_from_manager recognizes Azure SDK clients by type and ignores hosted_keys" + ) + + +async def test_custom_system_with_a_foreign_client_falls_back_to_the_environment( + monkeypatch: pytest.MonkeyPatch, ocr: Ocr +) -> None: + _configure_client( + monkeypatch, + client=_PlainSecretReader(), + system=KeyManagementSystem.CUSTOM, + key_management=KeyManagementSettings(access_mode="read_only"), + ) + + with _mistral_service() as server: + await ocr(server.base_url) + + assert server.requests[0].headers["authorization"] == "Bearer environment-key" + + +async def test_native_backend_supplies_ocr_credentials_without_a_python_reader( + monkeypatch: pytest.MonkeyPatch, rust_ocr: Ocr +) -> None: + from litellm.secret_managers import main as secret_manager_main + from litellm.secret_managers import secret_manager_handler + from litellm.secret_managers.aws_secret_manager_v2 import AWSSecretsManagerV2 + + def reject_python_read( + client: object, + key_manager: str, + secret_name: str, + key_management_settings: KeyManagementSettings | None = None, + ) -> str | None: + raise AssertionError("Rust must read the native backend directly") + + monkeypatch.setattr(secret_manager_handler, "get_secret_from_manager", reject_python_read) + monkeypatch.setattr(secret_manager_main, "get_secret_from_manager", reject_python_read) + with recording_service() as secrets, _mistral_service(expected_requests=2) as provider: + secrets.default_response = ResponseSpec(body={"SecretString": "native-key"}) + secrets.expected_requests = None + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "native-access") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "native-secret") + monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", secrets.base_url) + monkeypatch.delenv("MISTRAL_API_KEY", raising=False) + manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") + monkeypatch.setattr(litellm, "secret_manager_client", manager) + monkeypatch.setattr(litellm, "_key_management_system", KeyManagementSystem.AWS_SECRET_MANAGER) + monkeypatch.setattr(litellm, "_key_management_settings", KeyManagementSettings(hosted_keys=["MISTRAL_API_KEY"])) + monkeypatch.setattr(settings, "secret_manager", lambda: settings.SecretManager(readable=True, native=True)) + await rust_ocr(provider.base_url) + await rust_ocr(provider.base_url) + + assert len(secrets.requests) == 2, [(request.path, request.body) for request in secrets.requests] + assert all(request.headers["authorization"] == "Bearer native-key" for request in provider.requests) + assert all("Credential=native-access/" in request.headers["authorization"] for request in secrets.requests) + assert native._SecretManagerRuntime.from_client(manager) is getattr(manager, "_litellm_native_secret_manager") diff --git a/tests/test_litellm/rust_bridge/test_secret_manager.py b/tests/test_litellm/rust_bridge/test_secret_manager.py new file mode 100644 index 00000000000..e3e65a4e2bb --- /dev/null +++ b/tests/test_litellm/rust_bridge/test_secret_manager.py @@ -0,0 +1,1571 @@ +from __future__ import annotations + +import asyncio +import inspect +import json +import re +from collections.abc import Mapping +from dataclasses import dataclass +from functools import partial +from importlib import import_module +from types import SimpleNamespace +from typing import Final, Never +from urllib.parse import urlsplit + +import httpx +import pytest +from botocore.auth import SigV4Auth +from botocore.awsrequest import AWSRequest +from botocore.credentials import Credentials +from pydantic import JsonValue + +import litellm +from litellm.rust_bridge import bindings +from litellm.rust_bridge.bindings import NativeBinding +from litellm.rust_bridge.catalog import Rules, SecretManagerRule +from litellm.rust_bridge.configuration import Rollout +from litellm.rust_bridge.secret_manager import ( + NativeSecretManagerFactory, + NativeSecretManagerRuntime, + capture_secret_manager, + native_secret_manager_config, + resolve_native_provider_reader, + resolve_native_provider_writer, + resolve_native_secret_manager, +) +from litellm.secret_managers.aws_secret_manager_v2 import AWSSecretsManagerV2 +from litellm.secret_managers.cyberark_secret_manager import CyberArkSecretManager +from litellm.secret_managers.dispatch import get_secret_from_manager +from litellm.secret_managers.hashicorp_secret_manager import HashicorpSecretManager +from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem +from tests.test_litellm_rust.support.recording_server import ResponseSpec, recording_service + + +@pytest.fixture(autouse=True) +def preserve_manager_globals(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setattr(litellm, "secret_manager_client", litellm.secret_manager_client) + monkeypatch.setattr(litellm, "_key_management_system", litellm._key_management_system) + monkeypatch.setattr(litellm, "_key_management_settings", litellm._key_management_settings) + + +def _vault(monkeypatch: pytest.MonkeyPatch, address: str) -> HashicorpSecretManager: + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "premium_user", True) + monkeypatch.setenv("HCP_VAULT_ADDR", address) + monkeypatch.setenv("HCP_VAULT_TOKEN", "token") + return HashicorpSecretManager() + + +def _vault_body(value: str) -> dict[str, object]: + return { + "data": { + "data": {"key": value}, + "metadata": { + "created_time": "", + "deletion_time": "", + "custom_metadata": None, + "destroyed": False, + "version": 1, + }, + }, + "lease_id": "", + "lease_duration": 0, + "renewable": False, + "request_id": "", + "warnings": None, + "wrap_info": None, + } + + +@pytest.mark.parametrize("system", ("aws_secret_manager", "hashicorp_vault", "cyberark")) +async def test_python_native_handle_reuses_backend_across_sync_and_async_reads(system: str) -> None: + native: Final = pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + environment: Final = { + "AWS_REGION_NAME": "us-east-1", + "AWS_ACCESS_KEY_ID": "captured-access", + "AWS_SECRET_ACCESS_KEY": "captured-secret", + "AWS_BEDROCK_RUNTIME_ENDPOINT": server.base_url, + "AZURE_KEY_VAULT_URI": server.base_url, + "AZURE_AD_TOKEN": "azure-token", + "HCP_VAULT_ADDR": server.base_url, + "HCP_VAULT_TOKEN": "vault-token", + "CYBERARK_API_BASE": server.base_url, + "CYBERARK_API_KEY": "cyberark-key", + "CYBERARK_ACCOUNT": "account", + "CYBERARK_USERNAME": "reader", + } + responses: Final = { + "aws_secret_manager": {"SecretString": "native-value"}, + "azure_key_vault": {"value": "native-value"}, + "hashicorp_vault": _vault_body("native-value"), + "cyberark": "native-value", + } + server.default_response = ResponseSpec(body=responses[system]) + server.expected_requests = 2 if system == "cyberark" else 1 if system == "hashicorp_vault" else 3 + if system == "cyberark": + server.enqueue(ResponseSpec(body="authentication-token")) + handle: Final = native._SecretManagerRuntime.from_config(system, environment, enterprise_enabled=True) + expected: Final = '"native-value"' if system == "cyberark" else "native-value" + assert handle.read_secret("KEY") == expected + assert await handle.async_read_secret("KEY") == expected + assert handle.read_secret("KEY") == expected + + +def test_shared_initializer_captures_credentials_and_tracks_instance_settings(monkeypatch: pytest.MonkeyPatch) -> None: + native: Final = pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.default_response = ResponseSpec(body={"SecretString": "native-value"}) + server.expected_requests = 2 + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "captured-access") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "captured-secret") + monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", server.base_url) + from litellm.proxy.proxy_server import ProxyConfig + + monkeypatch.setattr(litellm, "_key_management_settings", KeyManagementSettings(aws_region_name="us-east-1")) + ProxyConfig().initialize_secret_manager(KeyManagementSystem.AWS_SECRET_MANAGER.value) + manager: Final = litellm.secret_manager_client + assert isinstance(manager, AWSSecretsManagerV2) + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "changed-access") + monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", "http://127.0.0.1:1") + first: Final = native._SecretManagerRuntime.from_client(manager) + assert first is not None + assert native._SecretManagerRuntime.from_client(manager) is first + assert first.read_secret("KEY") == "native-value" + manager.aws_region_name = "us-west-2" + second: Final = native._SecretManagerRuntime.from_client(manager) + assert second is not None + assert second is not first + assert second.read_secret("KEY") == "native-value" + assert "Credential=captured-access/" in server.requests[0].headers["authorization"] + assert "/us-east-1/" in server.requests[0].headers["authorization"] + assert "/us-west-2/" in server.requests[1].headers["authorization"] + + +def test_configuration_replacement_rebuilds_without_invalidating_existing_handles( + monkeypatch: pytest.MonkeyPatch, +) -> None: + native: Final = pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as first_server, recording_service() as second_server: + first_server.default_response = ResponseSpec(body=_vault_body("first")) + second_server.default_response = ResponseSpec(body=_vault_body("second")) + manager: Final = _vault(monkeypatch, first_server.base_url) + first: Final = native._SecretManagerRuntime.from_client(manager) + assert first is not None + assert first.read_secret("KEY") == "first" + manager.vault_addr = second_server.base_url + second: Final = native._SecretManagerRuntime.from_client(manager) + assert second is not None + assert second is not first + assert second.read_secret("KEY") == "second" + assert first.read_secret("KEY") == "first" + + +def test_custom_subclass_keeps_its_python_reader_under_native_selection() -> None: + class CustomManager(AWSSecretsManagerV2): + def sync_read_secret(self, secret_name: str, primary_secret_name: str | None = None) -> str: + return f"custom:{secret_name}" + + class DecliningFactory: + @staticmethod + def from_client(client: object) -> NativeSecretManagerRuntime | None: + assert native_secret_manager_config(client) is None + return None + + binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) + binding.override(DecliningFactory) + rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({"aws_secret_manager"})),) + assert ( + get_secret_from_manager( + CustomManager(aws_region_name="us-east-1"), "aws_secret_manager", "KEY", rules=rules, binding=binding + ) + == "custom:KEY" + ) + + +def test_read_dispatch_uses_native_backend_with_explicit_rules(monkeypatch: pytest.MonkeyPatch) -> None: + native: Final = pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.default_response = ResponseSpec(body=_vault_body("native-value")) + manager: Final = _vault(monkeypatch, server.base_url) + rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({"hashicorp_vault"})),) + assert get_secret_from_manager(manager, "hashicorp_vault", "KEY", rules=rules) == "native-value" + runtime: Final = native._SecretManagerRuntime.from_client(manager) + assert runtime is not None + assert runtime.read_secret("KEY") == "native-value" + assert len(server.requests) == 1 + + +@pytest.mark.parametrize("system", ("google_secret_manager", "hashicorp_vault", "cyberark")) +def test_enterprise_backends_cannot_initialize_without_entitlement(system: str) -> None: + native: Final = pytest.importorskip("litellm.rust_bridge._native") + with pytest.raises(ValueError, match=r"[Ee]nterprise|[Pp]remium"): + native._SecretManagerRuntime.from_config(system, {"CYBERARK_API_KEY": "key"}) + + +def test_azure_factory_rejects_unencrypted_vault_endpoints() -> None: + native: Final = pytest.importorskip("litellm.rust_bridge._native") + with pytest.raises(ValueError, match="https"): + native._SecretManagerRuntime.from_config("azure_key_vault", {"AZURE_KEY_VAULT_URI": "http://127.0.0.1:1"}) + + +def test_direct_builtin_constructor_can_use_native_without_registration(monkeypatch: pytest.MonkeyPatch) -> None: + native: Final = pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.default_response = ResponseSpec(body=_vault_body("native-value")) + manager: Final = _vault(monkeypatch, server.base_url) + handle: Final = native._SecretManagerRuntime.from_client(manager) + assert handle is not None + assert handle.read_secret("KEY") == "native-value" + + +def test_binding_rejects_a_different_backend_before_reading() -> None: + native: Final = pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.expected_requests = 0 + handle: Final = native._SecretManagerRuntime.from_config( + "hashicorp_vault", + {"HCP_VAULT_ADDR": server.base_url, "HCP_VAULT_TOKEN": "token"}, + enterprise_enabled=True, + ) + rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({"aws_secret_manager"})),) + with pytest.raises(ValueError, match="system does not match"): + resolve_native_secret_manager(handle, "aws_secret_manager", rules) + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_OPT_OUT)) +def test_python_selection_and_missing_extension_preserve_python_reader( + monkeypatch: pytest.MonkeyPatch, rollout: Rollout +) -> None: + binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) + binding.override(None) + with recording_service() as server: + server.default_response = ResponseSpec(body=_vault_body("python-value")) + manager: Final = _vault(monkeypatch, server.base_url) + rules: Final[Rules] = (SecretManagerRule(rollout, systems=frozenset({"hashicorp_vault"})),) + assert ( + get_secret_from_manager(manager, "hashicorp_vault", "KEY", rules=rules, binding=binding) == "python-value" + ) + assert server.requests[0].headers["x-vault-token"] == "token" + + +def test_required_native_missing_extension_does_not_read_python(monkeypatch: pytest.MonkeyPatch) -> None: + binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) + binding.override(None) + with recording_service() as server: + server.expected_requests = 0 + manager: Final = _vault(monkeypatch, server.base_url) + rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({"hashicorp_vault"})),) + with pytest.raises(RuntimeError, match="unavailable"): + get_secret_from_manager(manager, "hashicorp_vault", "KEY", rules=rules, binding=binding) + + +def test_native_failure_is_not_replayed_in_python(monkeypatch: pytest.MonkeyPatch) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.default_response = ResponseSpec(status=403, body={"errors": ["denied"]}) + manager: Final = _vault(monkeypatch, server.base_url) + rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_OPT_OUT, systems=frozenset({"hashicorp_vault"})),) + with pytest.raises(ValueError, match="HashiCorp Vault"): + get_secret_from_manager(manager, "hashicorp_vault", "KEY", rules=rules) + assert len(server.requests) == 1 + + +@pytest.mark.parametrize("explicit_capture", (False, True)) +def test_config_capture_preserves_credentials_and_excludes_unrelated_environment( + monkeypatch: pytest.MonkeyPatch, explicit_capture: bool +) -> None: + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "initial-access") + monkeypatch.setenv("SECRET_MANAGER_REFRESH_INTERVAL", "45") + monkeypatch.setenv("UNRELATED_PRIVATE_TOKEN", "unrelated-secret") + manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") + if explicit_capture: + capture_secret_manager(manager, "aws_secret_manager") + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "replacement-access") + captured: Final = native_secret_manager_config(manager) + assert captured is not None + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "replacement-access") + + retained: Final = native_secret_manager_config(manager) + + assert retained is captured + assert retained.system == "aws_secret_manager" + assert dict(retained.environment)["AWS_ACCESS_KEY_ID"] == "initial-access" + assert dict(retained.environment)["SECRET_MANAGER_REFRESH_INTERVAL"] == "45" + assert "UNRELATED_PRIVATE_TOKEN" not in dict(retained.environment) + assert "initial-access" not in repr(retained) + assert retained.settings == KeyManagementSettings().model_dump(mode="json") + + +def test_kms_sdk_client_capture_preserves_environment(monkeypatch: pytest.MonkeyPatch) -> None: + import boto3 + + monkeypatch.setenv("AWS_REGION_NAME", "captured-region") + client: Final = boto3.client( + "kms", region_name="us-east-1", aws_access_key_id="access", aws_secret_access_key="secret" + ) + try: + capture_secret_manager(client, "aws_kms") + monkeypatch.setenv("AWS_REGION_NAME", "replacement-region") + + captured: Final = native_secret_manager_config(client) + + assert captured is not None + assert captured.system == "aws_kms" + assert dict(captured.environment)["AWS_REGION_NAME"] == "captured-region" + finally: + client.close() + + +def test_same_named_custom_client_is_not_captured() -> None: + class AWSSecretsManagerV2: + pass + + client: Final = AWSSecretsManagerV2() + capture_secret_manager(client, "aws_secret_manager") + + assert native_secret_manager_config(client) is None + + +@dataclass(slots=True) +class _RecordingRuntime: + system: str + result: str | None + calls: tuple[tuple[str, Mapping[str, object] | None], ...] = () + + def read_secret(self, name: str, settings: Mapping[str, object] | None = None) -> str | None: + self.calls = (*self.calls, (name, settings)) + return self.result + + +@pytest.mark.parametrize("value", (None, " value\n")) +@pytest.mark.parametrize("settings", (None, KeyManagementSettings(primary_secret_name="primary"))) +def test_native_dispatch_forwards_settings_and_preserves_missing_or_unmodified_values( + value: str | None, settings: KeyManagementSettings | None +) -> None: + client: Final = object() + runtime: Final = _RecordingRuntime("aws_secret_manager", value) + + class Factory: + @staticmethod + def from_client(candidate: object) -> NativeSecretManagerRuntime: + assert candidate is client + return runtime + + binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) + binding.override(Factory) + rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({runtime.system})),) + + result: Final = get_secret_from_manager(client, runtime.system, "KEY", settings, rules=rules, binding=binding) + + assert result == value + assert runtime.calls == (("KEY", settings.model_dump(mode="json") if settings is not None else None),) + + +def test_native_system_mismatch_is_rejected_before_reading() -> None: + runtime: Final = _RecordingRuntime("hashicorp_vault", "wrong-provider") + + class Factory: + @staticmethod + def from_client(client: object) -> NativeSecretManagerRuntime: + return runtime + + binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) + binding.override(Factory) + rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({"aws_secret_manager"})),) + + with pytest.raises(ValueError, match="system does not match"): + get_secret_from_manager(object(), "aws_secret_manager", "KEY", rules=rules, binding=binding) + + assert runtime.calls == () + + +@pytest.mark.parametrize("system", ("custom", "local")) +def test_python_only_manager_types_never_construct_a_native_backend(system: str) -> None: + class ForbiddenFactory: + @staticmethod + def from_client(client: object) -> NativeSecretManagerRuntime: + raise AssertionError("custom and local clients cannot use native backends") + + binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) + binding.override(ForbiddenFactory) + rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({system})),) + + assert resolve_native_secret_manager(object(), system, rules, binding=binding) is None + + +@pytest.mark.parametrize("factory", (None, SimpleNamespace(from_client="not callable"))) +def test_invalid_native_factories_are_reported_as_unavailable(monkeypatch: pytest.MonkeyPatch, factory: object) -> None: + monkeypatch.setattr(bindings, "get_native_bridge", lambda: SimpleNamespace(_SecretManagerRuntime=factory)) + rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({"aws_secret_manager"})),) + + with pytest.raises(RuntimeError, match="runtime is unavailable"): + resolve_native_secret_manager(object(), "aws_secret_manager", rules) + + +def test_native_binding_accepts_a_callable_factory(monkeypatch: pytest.MonkeyPatch) -> None: + runtime: Final = _RecordingRuntime("aws_secret_manager", "value") + + class Factory: + @staticmethod + def from_client(client: object) -> NativeSecretManagerRuntime: + return runtime + + monkeypatch.setattr(bindings, "get_native_bridge", lambda: SimpleNamespace(_SecretManagerRuntime=Factory)) + rules: Final[Rules] = (SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({runtime.system})),) + + assert resolve_native_secret_manager(object(), runtime.system, rules) is runtime + + +@pytest.mark.parametrize("value", ("text", "", True, False, 42, 2**100, [1, "two"], {"nested": True}, None)) +async def test_aws_primary_values_match_python_handler(monkeypatch: pytest.MonkeyPatch, value: JsonValue) -> None: + native: Final = pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.default_response = ResponseSpec(body={"SecretString": json.dumps({"KEY": value})}) + server.expected_requests = 3 + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "test-access") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "test-secret") + monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", server.base_url) + manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") + settings: Final = KeyManagementSettings(primary_secret_name="primary") + reference: Final = get_secret_from_manager( + manager, "aws_secret_manager", "KEY", settings, rules=(SecretManagerRule(Rollout.PYTHON_ONLY),) + ) + actual: Final = get_secret_from_manager( + manager, "aws_secret_manager", "KEY", settings, rules=(SecretManagerRule(Rollout.RUST_REQUIRED),) + ) + handle: Final = native._SecretManagerRuntime.from_client(manager) + assert handle is not None + asynchronous: Final = await handle.read_secret_async("KEY", settings.model_dump(mode="json")) + assert type(actual) is type(reference) is type(value) + assert actual == reference == value + assert type(asynchronous) is type(reference) + assert asynchronous == reference + assert tuple(json.loads(request.raw_body) for request in server.requests) == ({"SecretId": "primary"},) * 3 + + +@pytest.mark.parametrize("primary", (None, "primary")) +@pytest.mark.parametrize( + ("status", "body"), + ( + (400, {"__type": "ResourceNotFoundException"}), + (403, {"__type": "AccessDeniedException"}), + (500, {"__type": "InternalServiceError"}), + (200, {"Name": "without-string"}), + (200, {"SecretString": ""}), + ), +) +def test_aws_absence_and_failed_reads_match_python_without_environment_fallback( + monkeypatch: pytest.MonkeyPatch, primary: str | None, status: int, body: dict[str, str] +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.default_response = ResponseSpec(status=status, body=body) + server.expected_requests = 2 + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "test-access") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "test-secret") + monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", server.base_url) + monkeypatch.setenv("KEY", "must-not-fall-back") + manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") + settings: Final = KeyManagementSettings(primary_secret_name=primary) + monkeypatch.setattr(litellm, "secret_manager_client", manager) + monkeypatch.setattr(litellm, "_key_management_system", KeyManagementSystem.AWS_SECRET_MANAGER) + monkeypatch.setattr(litellm, "_key_management_settings", settings) + main: Final = import_module("litellm.secret_managers.main") + monkeypatch.setattr( + main, + "get_secret_from_manager", + partial(get_secret_from_manager, rules=(SecretManagerRule(Rollout.PYTHON_ONLY),)), + ) + reference: Final = litellm.get_secret("KEY", "must-not-default") + monkeypatch.setattr( + main, + "get_secret_from_manager", + partial(get_secret_from_manager, rules=(SecretManagerRule(Rollout.RUST_REQUIRED),)), + ) + actual: Final = litellm.get_secret("KEY", "must-not-default") + assert actual == reference + assert actual == ("" if primary is None and body.get("SecretString") == "" else None) + + +@pytest.mark.parametrize("document", ("{", "not-json", "[1]", "null", "true", "42", '"text"')) +async def test_aws_primary_json_errors_preserve_python_exception_details( + monkeypatch: pytest.MonkeyPatch, document: str +) -> None: + native: Final = pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.default_response = ResponseSpec(body={"SecretString": document}) + server.expected_requests = 3 + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "test-access") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "test-secret") + monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", server.base_url) + manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") + settings: Final = KeyManagementSettings(primary_secret_name="primary") + with pytest.raises((json.JSONDecodeError, AttributeError)) as reference: + get_secret_from_manager( + manager, "aws_secret_manager", "KEY", settings, rules=(SecretManagerRule(Rollout.PYTHON_ONLY),) + ) + with pytest.raises(type(reference.value)) as actual: + get_secret_from_manager( + manager, "aws_secret_manager", "KEY", settings, rules=(SecretManagerRule(Rollout.RUST_REQUIRED),) + ) + assert actual.value.args == reference.value.args + if isinstance(reference.value, json.JSONDecodeError): + assert isinstance(actual.value, json.JSONDecodeError) + assert (actual.value.doc, actual.value.pos) == (reference.value.doc, reference.value.pos) + handle: Final = native._SecretManagerRuntime.from_client(manager) + assert handle is not None + with pytest.raises(type(reference.value)) as asynchronous: + await handle.read_secret_async("KEY", settings.model_dump(mode="json")) + assert asynchronous.value.args == reference.value.args + + +def _select_provider_reads(monkeypatch: pytest.MonkeyPatch, module_name: str, rollout: Rollout) -> None: + module: Final = import_module(module_name) + monkeypatch.setattr( + module, + "resolve_native_provider_reader", + partial(resolve_native_provider_reader, rules=(SecretManagerRule(rollout),)), + ) + if rollout is Rollout.RUST_REQUIRED: + monkeypatch.setattr(module, "_get_httpx_client", _forbid_python_http) + monkeypatch.setattr(module, "get_async_httpx_client", _forbid_python_http) + + +def _forbid_python_http(*args: object, **kwargs: object) -> Never: + raise AssertionError("native reads must not construct a Python HTTP client") + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +async def test_public_aws_reads_preserve_coroutines_and_per_call_credentials( + monkeypatch: pytest.MonkeyPatch, rollout: Rollout +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as initial, recording_service() as selected: + initial.expected_requests = 0 + selected.expected_requests = 2 + selected.default_response = ResponseSpec(body={"SecretString": "value"}) + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "environment-access") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "environment-secret") + monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", initial.base_url) + manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") + _select_provider_reads(monkeypatch, "litellm.secret_managers.aws_secret_manager_v2", rollout) + options: Final = { + "aws_region_name": "us-west-2", + "aws_bedrock_runtime_endpoint": selected.base_url, + "aws_access_key_id": "operation-access", + "aws_secret_access_key": "operation-secret", + "aws_session_token": "operation-session", + } + pending: Final = manager.async_read_secret(secret_name="KEY", optional_params=dict(options), timeout=2) + assert inspect.iscoroutine(pending) + assert selected.requests == [] + assert await asyncio.create_task(pending) == "value" + assert manager.sync_read_secret("KEY", dict(options), 2) == "value" + assert tuple(json.loads(request.raw_body) for request in selected.requests) == ({"SecretId": "KEY"},) * 2 + assert all( + "Credential=operation-access/" in request.headers["authorization"] + and "/us-west-2/" in request.headers["authorization"] + and request.headers["x-amz-security-token"] == "operation-session" + for request in selected.requests + ) + for request in selected.requests: + signed_headers: Final = request.headers["authorization"].split("SignedHeaders=")[1].split(",")[0].split(";") + signed_request: Final = AWSRequest( + method=request.method, + url=selected.base_url + request.path, + data=request.raw_body, + headers={name: request.headers[name] for name in signed_headers}, + ) + signed_request.context["timestamp"] = request.headers["x-amz-date"] + signer: Final = SigV4Auth( + Credentials( + options["aws_access_key_id"], options["aws_secret_access_key"], options["aws_session_token"] + ), + "secretsmanager", + options["aws_region_name"], + ) + string_to_sign: Final = signer.string_to_sign(signed_request, signer.canonical_request(signed_request)) + assert request.headers["authorization"].split("Signature=")[1] == signer.signature( + string_to_sign, signed_request + ) + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +async def test_public_aws_primary_reads_ignore_operation_overrides_like_python( + monkeypatch: pytest.MonkeyPatch, rollout: Rollout +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server, recording_service() as unused: + server.default_response = ResponseSpec(body={"SecretString": '{"KEY":true}'}) + server.expected_requests = 2 + unused.expected_requests = 0 + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "test-access") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "test-secret") + monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", server.base_url) + manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") + _select_provider_reads(monkeypatch, "litellm.secret_managers.aws_secret_manager_v2", rollout) + options: Final = {"aws_bedrock_runtime_endpoint": unused.base_url} + assert manager.sync_read_secret("KEY", options, 0, "primary") is True + assert ( + await manager.async_read_secret("KEY", optional_params=options, timeout=0, primary_secret_name="primary") + is True + ) + assert tuple(json.loads(request.raw_body) for request in server.requests) == ({"SecretId": "primary"},) * 2 + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +async def test_public_aws_bootstrap_names_only_bypass_sync_reads( + monkeypatch: pytest.MonkeyPatch, rollout: Rollout +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.default_response = ResponseSpec(body={"SecretString": "remote-access"}) + server.expected_requests = 1 + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "environment-access") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "environment-secret") + monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", server.base_url) + manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") + _select_provider_reads(monkeypatch, "litellm.secret_managers.aws_secret_manager_v2", rollout) + assert manager.sync_read_secret("AWS_ACCESS_KEY_ID") == "environment-access" + assert server.requests == [] + assert await manager.async_read_secret("AWS_ACCESS_KEY_ID") == "remote-access" + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +@pytest.mark.parametrize("timeout", (0.05, httpx.Timeout(1, read=0.05))) +async def test_public_aws_read_timeouts_follow_the_python_http_handler( + monkeypatch: pytest.MonkeyPatch, rollout: Rollout, timeout: float | httpx.Timeout +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.default_response = ResponseSpec(body={"SecretString": "too-late"}, delay=0.25) + server.expected_requests = 2 + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "test-access") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "test-secret") + monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", server.base_url) + manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") + _select_provider_reads(monkeypatch, "litellm.secret_managers.aws_secret_manager_v2", rollout) + assert manager.sync_read_secret("KEY", timeout=timeout) is None + assert await manager.async_read_secret("KEY", timeout=timeout) is None + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +async def test_public_vault_reads_keep_overrides_cache_and_coroutines( + monkeypatch: pytest.MonkeyPatch, rollout: Rollout +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.default_response = ResponseSpec(body=_vault_body("value")) + server.expected_requests = 2 + manager: Final = _vault(monkeypatch, server.base_url) + _select_provider_reads(monkeypatch, "litellm.secret_managers.hashicorp_secret_manager", rollout) + options: Final = {"secret_manager_settings": {"mount": "team", "path_prefix": "keys", "data": "key"}} + pending: Final = manager.async_read_secret("KEY", options) + assert inspect.iscoroutine(pending) + assert server.requests == [] + assert await asyncio.create_task(pending) == "value" + assert manager.sync_read_secret(secret_name="KEY", optional_params=options) == "value" + assert manager.sync_read_secret("KEY") == "value" + assert urlsplit(server.requests[0].path).path == "/v1/team/data/keys/KEY" + assert urlsplit(server.requests[1].path).path == "/v1/secret/data/KEY" + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +@pytest.mark.parametrize("status", (404, 403)) +async def test_public_vault_failed_reads_return_none_without_replay( + monkeypatch: pytest.MonkeyPatch, rollout: Rollout, status: int +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.default_response = ResponseSpec(status=status, body={"errors": ["unavailable"]}) + server.expected_requests = 2 + manager: Final = _vault(monkeypatch, server.base_url) + _select_provider_reads(monkeypatch, "litellm.secret_managers.hashicorp_secret_manager", rollout) + assert manager.sync_read_secret("KEY") is None + assert await manager.async_read_secret("KEY") is None + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +async def test_public_cyberark_reads_reuse_authentication_and_cached_values( + monkeypatch: pytest.MonkeyPatch, rollout: Rollout +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + from litellm.proxy import proxy_server + from litellm.secret_managers.cyberark_secret_manager import CyberArkSecretManager + + monkeypatch.setattr(proxy_server, "premium_user", True) + with recording_service() as server: + server.enqueue(ResponseSpec(body="authentication-token")) + server.default_response = ResponseSpec(body="secret-value") + server.expected_requests = 2 + monkeypatch.setenv("CYBERARK_API_BASE", server.base_url) + monkeypatch.setenv("CYBERARK_API_KEY", "api-key") + monkeypatch.setenv("CYBERARK_ACCOUNT", "account") + monkeypatch.setenv("CYBERARK_USERNAME", "reader") + manager: Final = CyberArkSecretManager() + _select_provider_reads(monkeypatch, "litellm.secret_managers.cyberark_secret_manager", rollout) + pending: Final = manager.async_read_secret(secret_name="KEY", timeout=0) + assert inspect.iscoroutine(pending) + assert server.requests == [] + assert await asyncio.create_task(pending) == '"secret-value"' + assert manager.sync_read_secret("KEY", timeout=0) == ( + '"secret-value"' if rollout is Rollout.RUST_REQUIRED else "secret-value" + ) + assert tuple(request.path for request in server.requests) == ( + "/authn/account/reader/authenticate", + "/secrets/account/variable/KEY", + ) + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +async def test_public_native_selection_and_missing_extension_keep_the_python_method( + monkeypatch: pytest.MonkeyPatch, rollout: Rollout +) -> None: + binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) + binding.override(None) + module: Final = import_module("litellm.secret_managers.aws_secret_manager_v2") + monkeypatch.setattr( + module, + "resolve_native_provider_reader", + partial(resolve_native_provider_reader, rules=(SecretManagerRule(rollout),), binding=binding), + ) + with recording_service() as server: + server.default_response = ResponseSpec(body={"SecretString": "python-value"}) + server.expected_requests = 0 if rollout is Rollout.RUST_REQUIRED else 1 + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "test-access") + monkeypatch.setenv("AWS_SECRET_ACCESS_KEY", "test-secret") + monkeypatch.setenv("AWS_BEDROCK_RUNTIME_ENDPOINT", server.base_url) + manager: Final = AWSSecretsManagerV2(aws_region_name="us-east-1") + pending: Final = manager.async_read_secret("KEY") + if rollout is Rollout.RUST_REQUIRED: + with pytest.raises(RuntimeError, match="runtime is unavailable"): + await pending + else: + assert await pending == "python-value" + + +def test_public_aws_bootstrap_read_does_not_initialize_a_backend(monkeypatch: pytest.MonkeyPatch) -> None: + module: Final = import_module("litellm.secret_managers.aws_secret_manager_v2") + monkeypatch.setattr(module, "resolve_native_provider_reader", _forbid_python_http) + monkeypatch.setenv("AWS_ACCESS_KEY_ID", "bootstrap-value") + assert AWSSecretsManagerV2().sync_read_secret("AWS_ACCESS_KEY_ID") == "bootstrap-value" + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +@pytest.mark.parametrize("override", ("", 42)) +def test_public_vault_prefix_overrides_match_python_string_conversion( + monkeypatch: pytest.MonkeyPatch, rollout: Rollout, override: str | int +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + monkeypatch.setenv("HCP_VAULT_PATH_PREFIX", "default-prefix") + with recording_service() as server: + server.default_response = ResponseSpec(body=_vault_body("value")) + server.expected_requests = 1 + manager: Final = _vault(monkeypatch, server.base_url) + _select_provider_reads(monkeypatch, "litellm.secret_managers.hashicorp_secret_manager", rollout) + assert manager.sync_read_secret("KEY", {"path_prefix": override}) == "value" + assert urlsplit(server.requests[0].path).path == ( + f"/v1/secret/data/{override}/KEY" if override else "/v1/secret/data/KEY" + ) + + +@pytest.mark.parametrize("value", ("native-value", None)) +def test_public_google_reader_uses_the_selected_binding_without_replaying_python( + monkeypatch: pytest.MonkeyPatch, value: str | None +) -> None: + from litellm.proxy import proxy_server + from litellm.secret_managers.google_secret_manager import GoogleSecretManager + + class Manager(GoogleSecretManager): + def sync_construct_request_headers(self) -> dict[str, str]: + raise AssertionError("selected native reads must not construct Python auth headers") + + class Reader(_RecordingRuntime): + def sync_read_secret( + self, + secret_name: str, + optional_params: Mapping[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> str | None: + return self.read_secret(secret_name) + + async def async_read_secret( + self, + secret_name: str, + optional_params: Mapping[str, object] | None = None, + timeout: float | httpx.Timeout | None = None, + ) -> str | None: + return self.sync_read_secret(secret_name) + + monkeypatch.setattr(proxy_server, "premium_user", True) + monkeypatch.setenv("GOOGLE_SECRET_MANAGER_PROJECT_ID", "project") + manager: Final = Manager() + runtime: Final = Reader("google_secret_manager", value) + + class Factory: + @staticmethod + def from_client(candidate: object) -> NativeSecretManagerRuntime: + assert candidate is manager + return runtime + + binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) + binding.override(Factory) + module: Final = import_module("litellm.secret_managers.google_secret_manager") + monkeypatch.setattr( + module, + "resolve_native_provider_reader", + partial(resolve_native_provider_reader, rules=(SecretManagerRule(Rollout.RUST_REQUIRED),), binding=binding), + ) + assert manager.get_secret_from_google_secret_manager(secret_name="KEY") == value + assert runtime.calls == (("KEY", None),) + + +def _cyberark(monkeypatch: pytest.MonkeyPatch, address: str) -> CyberArkSecretManager: + from litellm.proxy import proxy_server + + monkeypatch.setattr(proxy_server, "premium_user", True) + monkeypatch.setenv("CYBERARK_API_BASE", address) + monkeypatch.setenv("CYBERARK_API_KEY", "api-key") + monkeypatch.setenv("CYBERARK_ACCOUNT", "account") + monkeypatch.setenv("CYBERARK_USERNAME", "reader") + return CyberArkSecretManager() + + +def _select_cyberark_mutations(monkeypatch: pytest.MonkeyPatch, rollout: Rollout) -> None: + module: Final = import_module("litellm.secret_managers.cyberark_secret_manager") + _select_provider_reads(monkeypatch, module.__name__, rollout) + monkeypatch.setattr( + module, + "resolve_native_provider_writer", + partial(resolve_native_provider_writer, rules=(SecretManagerRule(rollout),)), + ) + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +async def test_public_cyberark_writes_and_deletes_share_the_read_cache( + monkeypatch: pytest.MonkeyPatch, + rollout: Rollout, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + for body in (b"token", b"old", {}, {}, b"provider-after-delete"): + server.enqueue(ResponseSpec(body=body)) + server.expected_requests = 5 + manager: Final = _cyberark(monkeypatch, server.base_url) + _select_cyberark_mutations(monkeypatch, rollout) + assert manager.sync_read_secret("KEY") == "old" + pending: Final = manager.async_write_secret( + "KEY", + "new-value", + "ignored", + {"ignored": object()}, + 0, + {"ignored": object()}, + ) + assert inspect.iscoroutine(pending) + assert len(server.requests) == 2 + assert await asyncio.create_task(pending) == { + "status": "success", + "message": "Secret KEY written successfully", + } + assert manager.sync_read_secret("KEY") == "new-value" + assert await manager.async_read_secret("KEY") == "new-value" + assert len(server.requests) == 4 + assert await manager.async_delete_secret(secret_name="KEY", recovery_window_in_days=None, timeout=0) == { + "status": "not_supported", + "message": "CyberArk Conjur does not support direct secret deletion. Use policy updates to remove variables.", + } + assert len(server.requests) == 4 + assert manager.sync_read_secret("KEY") == "provider-after-delete" + assert tuple(request.path for request in server.requests) == ( + "/authn/account/reader/authenticate", + "/secrets/account/variable/KEY", + "/policies/account/policy/root", + "/secrets/account/variable/KEY", + "/secrets/account/variable/KEY", + ) + assert server.requests[3].raw_body == b"new-value" + + +@pytest.mark.parametrize("status", (401, 403, 500)) +@pytest.mark.parametrize("authentication", (False, True)) +async def test_public_cyberark_write_errors_match_python_without_http_retries( + monkeypatch: pytest.MonkeyPatch, + status: int, + authentication: bool, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + responses: Final = ( + (ResponseSpec(status=status, body={}),) * 2 + if authentication + else ( + ResponseSpec(body=b"token"), + ResponseSpec(body={}), + ResponseSpec(status=status, body={}), + ) + ) + for response in responses * 2: + server.enqueue(response) + server.expected_requests = len(responses) * 2 + reference_manager: Final = _cyberark(monkeypatch, server.base_url) + _select_cyberark_mutations(monkeypatch, Rollout.PYTHON_ONLY) + reference: Final = await reference_manager.async_write_secret("KEY", "value") + native_manager: Final = _cyberark(monkeypatch, server.base_url) + _select_cyberark_mutations(monkeypatch, Rollout.RUST_REQUIRED) + actual: Final = await native_manager.async_write_secret("KEY", "value") + assert actual == reference + assert tuple(actual) == tuple(reference) + assert actual["status"] == "error" + assert str(status) in actual["message"] + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +async def test_public_cyberark_write_recovers_from_initial_policy_authentication_failure( + monkeypatch: pytest.MonkeyPatch, + rollout: Rollout, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.enqueue(ResponseSpec(status=401, body={})) + server.enqueue(ResponseSpec(body=b"token")) + server.enqueue(ResponseSpec(body={})) + server.expected_requests = 3 + manager: Final = _cyberark(monkeypatch, server.base_url) + _select_cyberark_mutations(monkeypatch, rollout) + assert await manager.async_write_secret("KEY", "value") == { + "status": "success", + "message": "Secret KEY written successfully", + } + assert tuple(request.path for request in server.requests) == ( + "/authn/account/reader/authenticate", + "/authn/account/reader/authenticate", + "/secrets/account/variable/KEY", + ) + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +@pytest.mark.parametrize("name", ("../KEY", "line\nKEY", "a\u2028b")) +async def test_public_cyberark_write_rejects_unsafe_names_before_authentication( + monkeypatch: pytest.MonkeyPatch, + rollout: Rollout, + name: str, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.expected_requests = 0 + manager: Final = _cyberark(monkeypatch, server.base_url) + _select_cyberark_mutations(monkeypatch, rollout) + assert await manager.async_write_secret(name, "value") == { + "status": "error", + "message": f"Invalid secret_name {name!r}", + } + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +@pytest.mark.parametrize("same_name", (False, True)) +async def test_public_cyberark_rotation_returns_the_write_response_and_retains_old_alias( + monkeypatch: pytest.MonkeyPatch, + rollout: Rollout, + same_name: bool, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + for body in (b"token", b"old-value", {}, {}): + server.enqueue(ResponseSpec(body=body)) + if rollout is Rollout.RUST_REQUIRED: + server.enqueue(ResponseSpec(body=b"new-value")) + server.expected_requests = 5 if rollout is Rollout.RUST_REQUIRED else 4 + manager: Final = _cyberark(monkeypatch, server.base_url) + _select_cyberark_mutations(monkeypatch, rollout) + new_name: Final = "OLD" if same_name else "NEW" + pending: Final = manager.async_rotate_secret("OLD", new_name, "new-value", {"ignored": object()}, 0) + assert inspect.iscoroutine(pending) + assert server.requests == [] + assert await asyncio.create_task(pending) == { + "status": "success", + "message": f"Secret {new_name} written successfully", + } + assert tuple(request.method for request in server.requests) == ( + ("POST", "GET", "POST", "POST", "GET") + if rollout is Rollout.RUST_REQUIRED + else ("POST", "GET", "POST", "POST") + ) + assert server.requests[3].raw_body == b"new-value" + + +@pytest.mark.parametrize("replacement", (None, b"wrong-value")) +async def test_public_cyberark_rotation_requires_a_fresh_matching_replacement( + monkeypatch: pytest.MonkeyPatch, + replacement: bytes | None, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + for body in (b"token", b"old-value", {}, {}): + server.enqueue(ResponseSpec(body=body)) + server.enqueue(ResponseSpec(status=404 if replacement is None else 200, body=replacement)) + server.expected_requests = 5 + manager: Final = _cyberark(monkeypatch, server.base_url) + _select_cyberark_mutations(monkeypatch, Rollout.RUST_REQUIRED) + message: Final = "Failed to verify new secret NEW" if replacement is None else "New secret value mismatch" + with pytest.raises(ValueError, match=message): + await manager.async_rotate_secret("OLD", "NEW", "new-value") + assert manager.sync_read_secret("OLD") == "old-value" + assert tuple(request.path for request in server.requests) == ( + "/authn/account/reader/authenticate", + "/secrets/account/variable/OLD", + "/policies/account/policy/root", + "/secrets/account/variable/NEW", + "/secrets/account/variable/NEW", + ) + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +async def test_public_cyberark_cached_authentication_does_not_retry_denied_reads( + monkeypatch: pytest.MonkeyPatch, + rollout: Rollout, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.enqueue(ResponseSpec(body=b"token")) + server.enqueue(ResponseSpec(body=b"value")) + server.enqueue(ResponseSpec(status=401, body={})) + server.expected_requests = 3 + manager: Final = _cyberark(monkeypatch, server.base_url) + _select_cyberark_mutations(monkeypatch, rollout) + assert manager.sync_read_secret("OLD") == "value" + assert await manager.async_read_secret("NEW") is None + + +async def test_cyberark_handler_errors_match_python_after_cached_authentication_is_denied( + monkeypatch: pytest.MonkeyPatch, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + for response in ( + ResponseSpec(body=b"token"), + ResponseSpec(body=b"value"), + ResponseSpec(status=401, body={}), + ) * 2: + server.enqueue(response) + server.expected_requests = 6 + reference_manager: Final = _cyberark(monkeypatch, server.base_url) + python_rules: Final = (SecretManagerRule(Rollout.PYTHON_ONLY),) + native_rules: Final = (SecretManagerRule(Rollout.RUST_REQUIRED),) + assert get_secret_from_manager(reference_manager, "cyberark", "OLD", rules=python_rules) == "value" + with pytest.raises(ValueError, match="No secret found in CyberArk Secret Manager for NEW") as reference: + get_secret_from_manager(reference_manager, "cyberark", "NEW", rules=python_rules) + native_manager: Final = _cyberark(monkeypatch, server.base_url) + _select_cyberark_mutations(monkeypatch, Rollout.RUST_REQUIRED) + assert get_secret_from_manager(native_manager, "cyberark", "OLD", rules=native_rules) == "value" + with pytest.raises(ValueError, match="No secret found in CyberArk Secret Manager for NEW") as actual: + get_secret_from_manager(native_manager, "cyberark", "NEW", rules=native_rules) + assert actual.value.args == reference.value.args + + +async def test_public_cyberark_connection_errors_match_python( + monkeypatch: pytest.MonkeyPatch, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.expected_requests = 0 + address: Final = server.base_url + reference_manager: Final = _cyberark(monkeypatch, address) + _select_cyberark_mutations(monkeypatch, Rollout.PYTHON_ONLY) + reference: Final = await reference_manager.async_write_secret("KEY", "value") + native_manager: Final = _cyberark(monkeypatch, address) + _select_cyberark_mutations(monkeypatch, Rollout.RUST_REQUIRED) + actual: Final = await native_manager.async_write_secret("KEY", "value") + assert actual == reference + assert actual["status"] == "error" + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +@pytest.mark.parametrize("operation", ("write", "delete", "rotate")) +async def test_public_cyberark_mutations_preserve_missing_extension_selection( + monkeypatch: pytest.MonkeyPatch, + rollout: Rollout, + operation: str, +) -> None: + binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) + binding.override(None) + module: Final = import_module("litellm.secret_managers.cyberark_secret_manager") + _select_provider_reads(monkeypatch, module.__name__, Rollout.PYTHON_ONLY) + monkeypatch.setattr( + module, + "resolve_native_provider_writer", + partial(resolve_native_provider_writer, rules=(SecretManagerRule(rollout),), binding=binding), + ) + with recording_service() as server: + bodies: Final = ( + () + if rollout is Rollout.RUST_REQUIRED or operation == "delete" + else ((b"token", b"old", {}, {}) if operation == "rotate" else (b"token", {}, {})) + ) + for body in bodies: + server.enqueue(ResponseSpec(body=body)) + server.expected_requests = len(bodies) + manager: Final = _cyberark(monkeypatch, server.base_url) + call: Final = { + "write": partial(manager.async_write_secret, "KEY", "value"), + "delete": partial(manager.async_delete_secret, "KEY"), + "rotate": partial(manager.async_rotate_secret, "OLD", "NEW", "value"), + }[operation] + pending: Final = call() + assert inspect.iscoroutine(pending) + assert server.requests == [] + if rollout is Rollout.RUST_REQUIRED: + with pytest.raises(RuntimeError, match="runtime is unavailable"): + await pending + else: + result: Final = await pending + assert result["status"] == ("not_supported" if operation == "delete" else "success") + + +async def test_public_cyberark_rotation_stops_after_a_failed_write(monkeypatch: pytest.MonkeyPatch) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + for body in (b"token", b"old-value", {}): + server.enqueue(ResponseSpec(body=body)) + server.enqueue(ResponseSpec(status=401, body={})) + server.expected_requests = 4 + manager: Final = _cyberark(monkeypatch, server.base_url) + _select_cyberark_mutations(monkeypatch, Rollout.RUST_REQUIRED) + response: Final = await manager.async_rotate_secret("OLD", "NEW", "new-value") + assert response["status"] == "error" + assert "401" in response["message"] + assert manager.sync_read_secret("OLD") == "old-value" + assert tuple(request.method for request in server.requests) == ("POST", "GET", "POST", "POST") + + +def _select_vault_mutations(monkeypatch: pytest.MonkeyPatch, rollout: Rollout) -> None: + module: Final = import_module("litellm.secret_managers.hashicorp_secret_manager") + _select_provider_reads(monkeypatch, module.__name__, rollout) + monkeypatch.setattr( + module, + "resolve_native_provider_writer", + partial(resolve_native_provider_writer, rules=(SecretManagerRule(rollout),)), + ) + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +@pytest.mark.parametrize("description", (None, "", "purpose")) +async def test_public_vault_writes_preserve_complete_responses_and_request_fields( + monkeypatch: pytest.MonkeyPatch, + rollout: Rollout, + description: str | None, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + response: Final = { + "request_id": "test-request", + "data": {"version": 2, "custom_metadata": {"large": 2**100}}, + "warnings": ["test-warning"], + "unknown_field": {"nested": [None, True, ""]}, + } + server.enqueue(ResponseSpec(body=response)) + server.expected_requests = 1 + manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, rollout) + options: Final = { + "secret_manager_settings": {"namespace": "team", "mount": "kv", "path_prefix": "app", "data": "token"} + } + pending: Final = manager.async_write_secret("KEY", "value", description, options, 2, {"ignored": object()}) + assert inspect.iscoroutine(pending) + assert server.requests == [] + result: Final = await asyncio.create_task(pending) + assert result == response + assert tuple(result) == tuple(response) + request: Final = server.requests[0] + assert request.method == "POST" + assert request.headers["x-vault-token"] == "token" + namespace: Final = request.headers.get("x-vault-namespace") + assert request.path == ("/v1/kv/data/app/KEY" if namespace else "/v1/team/kv/data/app/KEY") + assert namespace in (None, "team") + assert json.loads(request.raw_body) == { + "data": {"token": "value", **({"description": description} if description else {})}, + } + assert options == { + "secret_manager_settings": {"namespace": "team", "mount": "kv", "path_prefix": "app", "data": "token"} + } + + +@pytest.mark.parametrize("operation", ("write", "delete")) +@pytest.mark.parametrize("status", (400, 403, 500)) +async def test_public_vault_mutation_http_errors_match_python_without_retry( + monkeypatch: pytest.MonkeyPatch, + operation: str, + status: int, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.default_response = ResponseSpec(status=status, body={"errors": ["denied"]}) + server.expected_requests = 2 + options: Final = {"namespace": "team", "mount": "kv", "path_prefix": "prefix"} + reference_manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.PYTHON_ONLY) + reference: Final = ( + await reference_manager.async_write_secret("KEY", "value", optional_params=options) + if operation == "write" + else await reference_manager.async_delete_secret("KEY", optional_params=options) + ) + native_manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) + actual: Final = ( + await native_manager.async_write_secret("KEY", "value", optional_params=options) + if operation == "write" + else await native_manager.async_delete_secret("KEY", optional_params=options) + ) + assert actual == reference + assert tuple(actual) == tuple(reference) + assert actual["status"] == "error" + assert str(status) in actual["message"] + + +@pytest.mark.parametrize("body", (b"{", b"", b"null", b"[1,2]", b'{"large":1267650600228229401496703205376}')) +async def test_public_vault_write_response_conversion_matches_python( + monkeypatch: pytest.MonkeyPatch, + body: bytes, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.default_response = ResponseSpec(body=body) + server.expected_requests = 2 + reference_manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.PYTHON_ONLY) + reference: Final = await reference_manager.async_write_secret("KEY", "value") + native_manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) + actual: Final = await native_manager.async_write_secret("KEY", "value") + assert type(actual) is type(reference) + assert actual == reference + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +async def test_public_vault_deletion_invalidates_cached_fields( + monkeypatch: pytest.MonkeyPatch, + rollout: Rollout, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.enqueue(ResponseSpec(body=_vault_body("old"))) + server.enqueue(ResponseSpec(status=204, body=b"")) + server.enqueue(ResponseSpec(body=_vault_body("after-delete"))) + server.expected_requests = 3 + manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, rollout) + assert manager.sync_read_secret("KEY") == "old" + pending: Final = manager.async_delete_secret("KEY", None, {"ignored": object()}, 2) + assert inspect.iscoroutine(pending) + assert len(server.requests) == 1 + assert await asyncio.create_task(pending) == {"status": "success", "message": "Secret KEY deleted successfully"} + assert await manager.async_read_secret("KEY") == "after-delete" + assert tuple(request.method for request in server.requests) == ("GET", "DELETE", "GET") + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +@pytest.mark.parametrize("same_name", (False, True)) +@pytest.mark.parametrize("delete_status", (204, 403)) +async def test_public_vault_rotation_preserves_response_and_best_effort_deletion( + monkeypatch: pytest.MonkeyPatch, + rollout: Rollout, + same_name: bool, + delete_status: int, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + response: Final = {"request_id": "write-id", "data": {"version": 3}, "extra": [1, 2]} + server.enqueue(ResponseSpec(body=b"current-existence-is-status-only")) + server.enqueue(ResponseSpec(body=response)) + server.enqueue(ResponseSpec(body=_vault_body("replacement"))) + if not same_name: + server.enqueue(ResponseSpec(status=delete_status, body=b"")) + server.expected_requests = 3 if same_name else 4 + manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, rollout) + new_name: Final = "OLD" if same_name else "NEW" + pending: Final = manager.async_rotate_secret("OLD", new_name, "replacement", timeout=2) + assert inspect.iscoroutine(pending) + assert server.requests == [] + assert await asyncio.create_task(pending) == response + assert tuple(request.method for request in server.requests) == ( + ("GET", "POST", "GET") if same_name else ("GET", "POST", "GET", "DELETE") + ) + assert json.loads(server.requests[1].raw_body) == { + "data": {"key": "replacement", "description": "Rotated from OLD"}, + } + assert urlsplit(server.requests[2].path).path == f"/v1/secret/data/{new_name}" + + +@pytest.mark.parametrize("stage", ("current", "write", "verify")) +@pytest.mark.parametrize("status", (404, 403, 500)) +async def test_public_vault_rotation_failure_messages_and_request_counts_match_python( + monkeypatch: pytest.MonkeyPatch, + stage: str, + status: int, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + before: Final = ( + () + if stage == "current" + else ( + (ResponseSpec(body=_vault_body("old")),) + if stage == "write" + else ( + ResponseSpec(body=_vault_body("old")), + ResponseSpec(body={"data": {"version": 2}}), + ) + ) + ) + responses: Final = (*before, ResponseSpec(status=status, body={"errors": ["denied"]})) + for response in responses * 2: + server.enqueue(response) + server.expected_requests = len(responses) * 2 + reference_manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.PYTHON_ONLY) + reference: Final = await reference_manager.async_rotate_secret("OLD", "NEW", "value") + native_manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) + actual: Final = await native_manager.async_rotate_secret("OLD", "NEW", "value") + assert actual == reference + assert actual["status"] == "error" + expected_paths: Final = ( + ("/v1/secret/data/OLD",) + + (("/v1/secret/data/NEW",) if stage != "current" else ()) + + (("/v1/secret/data/NEW",) if stage == "verify" else ()) + ) + assert tuple(urlsplit(request.path).path for request in server.requests) == expected_paths * 2 + + +@pytest.mark.parametrize("value", (None, "different", True, 42, 2**100, [1, "two"], [2**100], {"nested": 2**100})) +async def test_public_vault_rotation_mismatches_do_not_delete_the_old_alias( + monkeypatch: pytest.MonkeyPatch, + value: JsonValue, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + responses: Final = ( + ResponseSpec(body=_vault_body("old")), + ResponseSpec(body={"data": {"version": 2}}), + ResponseSpec(body={"data": {"data": {"key": value}}}), + ) + for response in responses * 2: + server.enqueue(response) + server.expected_requests = 6 + reference_manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.PYTHON_ONLY) + reference: Final = await reference_manager.async_rotate_secret("OLD", "NEW", "value") + native_manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) + actual: Final = await native_manager.async_rotate_secret("OLD", "NEW", "value") + assert actual == reference + assert actual["status"] == "error" + assert all(request.method != "DELETE" for request in server.requests) + + +@pytest.mark.parametrize("operation", ("write", "delete", "rotate")) +async def test_public_vault_mutation_timeouts_match_python( + monkeypatch: pytest.MonkeyPatch, + operation: str, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.default_response = ResponseSpec(body=_vault_body("value"), delay=0.25) + server.expected_requests = 2 + reference_manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.PYTHON_ONLY) + reference: Final = await { + "write": partial(reference_manager.async_write_secret, "KEY", "value"), + "delete": partial(reference_manager.async_delete_secret, "KEY"), + "rotate": partial(reference_manager.async_rotate_secret, "OLD", "NEW", "value"), + }[operation](timeout=0.05) + native_manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) + actual: Final = await { + "write": partial(native_manager.async_write_secret, "KEY", "value"), + "delete": partial(native_manager.async_delete_secret, "KEY"), + "rotate": partial(native_manager.async_rotate_secret, "OLD", "NEW", "value"), + }[operation](timeout=0.05) + if operation == "write": + assert isinstance(actual["message"], str) + assert isinstance(reference["message"], str) + pattern: Final = r"time taken=(\d+(?:\.\d+)?) seconds" + assert re.sub(pattern, "time taken= seconds", actual["message"]) == re.sub( + pattern, + "time taken= seconds", + reference["message"], + ) + elapsed: Final = re.search(pattern, actual["message"]) + assert elapsed is not None + assert float(elapsed[1]) >= 0.05 + else: + assert actual == reference + assert actual["status"] == "error" + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +@pytest.mark.parametrize("operation", ("write", "delete", "rotate")) +async def test_public_vault_unsafe_names_fail_before_authentication( + monkeypatch: pytest.MonkeyPatch, + rollout: Rollout, + operation: str, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.expected_requests = 0 + manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, rollout) + result: Final = await { + "write": partial(manager.async_write_secret, "../KEY", "value"), + "delete": partial(manager.async_delete_secret, "../KEY"), + "rotate": partial(manager.async_rotate_secret, "../KEY", "NEW", "value"), + }[operation]() + assert result == {"status": "error", "message": "Invalid secret_name '../KEY'"} + + +@pytest.mark.parametrize("operation", ("write", "delete", "rotate")) +async def test_public_vault_authentication_errors_match_python( + monkeypatch: pytest.MonkeyPatch, + operation: str, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + monkeypatch.setenv("HCP_VAULT_APPROLE_ROLE_ID", "role") + monkeypatch.setenv("HCP_VAULT_APPROLE_SECRET_ID", "secret-id") + with recording_service() as server: + server.default_response = ResponseSpec(status=403, body={"errors": ["denied"]}) + server.expected_requests = 2 + reference_manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.PYTHON_ONLY) + reference: Final = await { + "write": partial(reference_manager.async_write_secret, "KEY", "value"), + "delete": partial(reference_manager.async_delete_secret, "KEY"), + "rotate": partial(reference_manager.async_rotate_secret, "OLD", "NEW", "value"), + }[operation]() + native_manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) + actual: Final = await { + "write": partial(native_manager.async_write_secret, "KEY", "value"), + "delete": partial(native_manager.async_delete_secret, "KEY"), + "rotate": partial(native_manager.async_rotate_secret, "OLD", "NEW", "value"), + }[operation]() + assert actual == reference + assert actual["status"] == "error" + assert tuple(request.path for request in server.requests) == ("/v1/auth/approle/login",) * 2 + + +async def test_public_vault_rotation_stops_on_a_success_response_containing_an_error( + monkeypatch: pytest.MonkeyPatch, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + result: Final = {"status": "error", "message": "write rejected", "extra": 2**100} + for response in (ResponseSpec(body=_vault_body("old")), ResponseSpec(body=result)) * 2: + server.enqueue(response) + server.expected_requests = 4 + reference_manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.PYTHON_ONLY) + reference: Final = await reference_manager.async_rotate_secret("OLD", "NEW", "value") + native_manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) + actual: Final = await native_manager.async_rotate_secret("OLD", "NEW", "value") + assert actual == reference == result + assert tuple(request.method for request in server.requests) == ("GET", "POST", "GET", "POST") + + +async def test_public_native_vault_write_invalidates_stale_cached_values(monkeypatch: pytest.MonkeyPatch) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.enqueue(ResponseSpec(body=_vault_body("old"))) + server.enqueue(ResponseSpec(body={"data": {"version": 2}})) + server.enqueue(ResponseSpec(body=_vault_body("new"))) + server.expected_requests = 3 + manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) + assert manager.sync_read_secret("KEY") == "old" + assert await manager.async_write_secret("KEY", "new") == {"data": {"version": 2}} + assert manager.sync_read_secret("KEY") == "new" + assert tuple(request.method for request in server.requests) == ("GET", "POST", "GET") + + +async def test_public_native_vault_write_rejects_description_overwriting_the_secret( + monkeypatch: pytest.MonkeyPatch, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + server.expected_requests = 0 + manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) + assert await manager.async_write_secret("KEY", "value", "description", {"data": "description"}) == { + "status": "error", + "message": "HashiCorp Vault data key conflicts with description", + } + + +@pytest.mark.parametrize("rollout", (Rollout.PYTHON_ONLY, Rollout.RUST_REQUIRED)) +@pytest.mark.parametrize("operation", ("write", "delete", "rotate")) +async def test_public_vault_mutations_preserve_missing_extension_selection( + monkeypatch: pytest.MonkeyPatch, + rollout: Rollout, + operation: str, +) -> None: + binding: Final[NativeBinding[NativeSecretManagerFactory]] = NativeBinding("unused", validate=lambda value: None) + binding.override(None) + module: Final = import_module("litellm.secret_managers.hashicorp_secret_manager") + _select_provider_reads(monkeypatch, module.__name__, Rollout.PYTHON_ONLY) + monkeypatch.setattr( + module, + "resolve_native_provider_writer", + partial(resolve_native_provider_writer, rules=(SecretManagerRule(rollout),), binding=binding), + ) + with recording_service() as server: + server.default_response = ResponseSpec(body=_vault_body("value")) + server.expected_requests = 0 if rollout is Rollout.RUST_REQUIRED else (4 if operation == "rotate" else 1) + manager: Final = _vault(monkeypatch, server.base_url) + pending: Final = { + "write": partial(manager.async_write_secret, "KEY", "value"), + "delete": partial(manager.async_delete_secret, "KEY"), + "rotate": partial(manager.async_rotate_secret, "OLD", "NEW", "value"), + }[operation]() + assert inspect.iscoroutine(pending) + assert server.requests == [] + if rollout is Rollout.RUST_REQUIRED: + with pytest.raises(RuntimeError, match="runtime is unavailable"): + await pending + else: + result: Final = await pending + assert result == ( + {"status": "success", "message": "Secret KEY deleted successfully"} + if operation == "delete" + else _vault_body("value") + ) + + +@pytest.mark.parametrize( + "body", (None, [], 42, 2**100, {"data": 2**100}, {"data": None}, {"data": {"data": []}}, {}, {"data": {}}) +) +async def test_public_vault_rotation_preserves_malformed_verification_errors( + monkeypatch: pytest.MonkeyPatch, + body: JsonValue, +) -> None: + pytest.importorskip("litellm.rust_bridge._native") + with recording_service() as server: + responses: Final = ( + ResponseSpec(body=_vault_body("old")), + ResponseSpec(body={"data": {"version": 2}}), + ResponseSpec(body=body), + ) + for response in responses * 2: + server.enqueue(response) + server.expected_requests = 6 + reference_manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.PYTHON_ONLY) + reference: Final = await reference_manager.async_rotate_secret("OLD", "NEW", "value") + native_manager: Final = _vault(monkeypatch, server.base_url) + _select_vault_mutations(monkeypatch, Rollout.RUST_REQUIRED) + actual: Final = await native_manager.async_rotate_secret("OLD", "NEW", "value") + assert actual == reference + assert actual["status"] == "error" + assert all(request.method != "DELETE" for request in server.requests) diff --git a/tests/test_litellm/rust_bridge/test_settings.py b/tests/test_litellm/rust_bridge/test_settings.py index 216961ab640..7650195b3c5 100644 --- a/tests/test_litellm/rust_bridge/test_settings.py +++ b/tests/test_litellm/rust_bridge/test_settings.py @@ -6,7 +6,9 @@ import pytest import litellm from litellm.integrations.custom_secret_manager import CustomSecretManager from litellm.llms.custom_httpx.http_handler import default_user_agent -from litellm.rust_bridge import settings +from litellm.rust_bridge import catalog, settings +from litellm.rust_bridge.catalog import SecretManagerRule +from litellm.rust_bridge.configuration import Rollout from litellm.secret_managers.main import get_secret_str from litellm.types.secret_managers.main import KeyManagementSettings, KeyManagementSystem @@ -79,6 +81,11 @@ class _VaultSecrets(CustomSecretManager): return self.secrets.get(secret_name) +_RUST_FOR_CUSTOM: Final = ( + SecretManagerRule(Rollout.RUST_REQUIRED, systems=frozenset({KeyManagementSystem.CUSTOM.value})), +) + + @pytest.mark.parametrize( ("access_mode", "readable"), [("read_only", True), ("read_and_write", True), ("write_only", False)], @@ -91,14 +98,46 @@ def test_secret_manager_is_readable_only_when_litellm_would_read_secrets_from_it monkeypatch.setattr(litellm, "_key_management_system", KeyManagementSystem.CUSTOM) monkeypatch.setattr(litellm, "_key_management_settings", KeyManagementSettings(access_mode=access_mode)) - assert settings.secret_manager() == settings.SecretManager(readable=readable) + assert settings.secret_manager(rules=()) == settings.SecretManager(readable=readable, native=False) assert (get_secret_str("MISTRAL_API_KEY") == "vault-key") is readable +@pytest.mark.parametrize( + ("system", "access_mode", "rules", "native"), + [ + (KeyManagementSystem.CUSTOM, "read_only", _RUST_FOR_CUSTOM, True), + (KeyManagementSystem.CUSTOM, "read_only", (), False), + ( + KeyManagementSystem.CUSTOM, + "read_only", + (SecretManagerRule(Rollout.PYTHON_ONLY, systems=frozenset({KeyManagementSystem.CUSTOM.value})),), + False, + ), + (KeyManagementSystem.CUSTOM, "write_only", _RUST_FOR_CUSTOM, False), + (None, "read_only", _RUST_FOR_CUSTOM, False), + (KeyManagementSystem.AWS_SECRET_MANAGER, "read_only", _RUST_FOR_CUSTOM, False), + ], +) +def test_secret_manager_is_native_only_when_the_rules_select_rust_for_its_system( + monkeypatch: pytest.MonkeyPatch, + system: KeyManagementSystem | None, + access_mode: str, + rules: catalog.Rules, + native: bool, +) -> None: + monkeypatch.setattr(litellm, "secret_manager_client", _VaultSecrets({})) + monkeypatch.setattr(litellm, "_key_management_system", system) + monkeypatch.setattr(litellm, "_key_management_settings", KeyManagementSettings(access_mode=access_mode)) + + assert settings.secret_manager(rules=rules) == settings.SecretManager( + readable=access_mode != "write_only", native=native + ) + + def test_secret_manager_is_not_readable_without_a_client(monkeypatch: pytest.MonkeyPatch) -> None: monkeypatch.setattr(litellm, "secret_manager_client", None) - assert settings.secret_manager() == settings.SecretManager(readable=False) + assert settings.secret_manager(rules=_RUST_FOR_CUSTOM) == settings.SecretManager(readable=False, native=False) def test_secret_manager_projects_custom_settings(monkeypatch: pytest.MonkeyPatch) -> None: diff --git a/tests/test_litellm_rust/support/recording_server.py b/tests/test_litellm_rust/support/recording_server.py index 3eea47751d3..74ca2cda1c5 100644 --- a/tests/test_litellm_rust/support/recording_server.py +++ b/tests/test_litellm_rust/support/recording_server.py @@ -28,6 +28,8 @@ class ResponseSpec: events: tuple[tuple[str, object], ...] = () def payloads(self) -> tuple[bytes, ...]: + if isinstance(self.body, bytes): + return (self.body,) if not self.events: return (json.dumps(self.body).encode(),) return tuple(f"event: {event}\ndata: {json.dumps(data)}\n\n".encode() for event, data in self.events) @@ -95,6 +97,7 @@ def recording_service() -> Iterator[RecordingServer]: do_POST = _handle do_GET = _handle + do_DELETE = _handle def log_message(self, format: str, *args: object) -> None: pass From cddc53464e0b964637123be0821eb6c623903991 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 08:44:59 -0700 Subject: [PATCH 06/10] test: deflake fuzzy picker, breached-password HIBP, and MCP stdio timeout tests (rolling deflake 2026-09-22) (#42125) * test(autoroute): wait for a valid fuzzy selection index and cancel the prompt on driver failure Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(autoroute): read the fuzzy selection through the public InquirerPy property Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): inject the HIBP client into change_password so the breached-password test never touches the network Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(mcp): only use the 200ms read timeout in the silent mode of the transport completion test Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(proxy): record HIBP requests so the ordering test asserts no lookup happened Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * ci(codeql): filter the weak-sensitive-data-hashing false positive on the HIBP k-anonymity lookup Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .github/workflows/codeql.yml | 8 ++- litellm/proxy/auth/password_policy.py | 4 +- .../password_endpoints.py | 10 +++- .../test_mcp_client.py | 5 +- .../proxy/client/cli/autoroute/test_wizard.py | 43 ++++++++++---- .../test_password_endpoints.py | 58 +++++++++++++------ 6 files changed, 92 insertions(+), 36 deletions(-) diff --git a/.github/workflows/codeql.yml b/.github/workflows/codeql.yml index d3a165a11da..9a85ced57f6 100644 --- a/.github/workflows/codeql.yml +++ b/.github/workflows/codeql.yml @@ -67,12 +67,18 @@ jobs: # further up the stack are modified. The suppression is scoped to this one # file/rule pair via SARIF post-filtering so every other callsite of # py/weak-sensitive-data-hashing in the repository continues to be analyzed. - - name: Filter SARIF (OCI sha256) + # The same query fires on the HIBP k-anonymity lookup in + # litellm/proxy/auth/password_policy.py, where the password's SHA-1 is only + # a lookup key into the haveibeenpwned range API (the protocol mandates + # SHA-1) and the digest itself never leaves the proxy beyond its first 5 + # characters. + - name: Filter SARIF (OCI sha256, HIBP sha1) if: matrix.language == 'python' uses: advanced-security/filter-sarif@2da736ff05ef065cb2894ac6892e47b5eac2c3c0 # v1.1 with: patterns: | -litellm/llms/oci/common_utils.py:py/weak-sensitive-data-hashing + -litellm/proxy/auth/password_policy.py:py/weak-sensitive-data-hashing input: sarif-results/python.sarif output: sarif-results/python.sarif diff --git a/litellm/proxy/auth/password_policy.py b/litellm/proxy/auth/password_policy.py index 7f06a0993d3..a883cfd6f35 100644 --- a/litellm/proxy/auth/password_policy.py +++ b/litellm/proxy/auth/password_policy.py @@ -107,7 +107,7 @@ def validate_password_policy(password: str, general_settings: Mapping[str, objec ) -def _hibp_client() -> AsyncHTTPHandler: +def get_hibp_client() -> AsyncHTTPHandler: return get_async_httpx_client( llm_provider=httpxSpecialProvider.PasswordBreachCheck, params={"timeout": HIBP_TIMEOUT_SECONDS}, # mutable-ok: callee takes a bare dict (PEP 589) @@ -155,7 +155,7 @@ async def is_password_breached( corpus, or HIBP is unreachable (fail open).""" if not is_breach_check_enabled(general_settings): return False - return await _is_password_breached(password, client if client is not None else _hibp_client()) + return await _is_password_breached(password, client if client is not None else get_hibp_client()) def breached_password_error() -> ProxyException: diff --git a/litellm/proxy/management_endpoints/password_endpoints.py b/litellm/proxy/management_endpoints/password_endpoints.py index 99b7c994b40..16f3e7dfcfb 100644 --- a/litellm/proxy/management_endpoints/password_endpoints.py +++ b/litellm/proxy/management_endpoints/password_endpoints.py @@ -14,6 +14,7 @@ from fastapi import APIRouter, Depends, HTTPException from pydantic import TypeAdapter from litellm._logging import verbose_proxy_logger +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.proxy._types import ( UI_TEAM_ID, ChangePasswordRequest, @@ -24,7 +25,11 @@ from litellm.proxy._types import ( UserAPIKeyAuth, ) from litellm.proxy.auth.login_utils import PASSWORD_SESSION_METADATA -from litellm.proxy.auth.password_policy import validate_password_not_breached, validate_password_policy +from litellm.proxy.auth.password_policy import ( + get_hibp_client, + validate_password_not_breached, + validate_password_policy, +) from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.management_endpoints.session_endpoints import revoke_ui_session_keys from litellm.proxy.management_helpers.audit_logs import create_object_audit_log @@ -71,6 +76,7 @@ def _user_table( async def change_password( data: ChangePasswordRequest, user_api_key_dict: Annotated[UserAPIKeyAuth, Depends(user_api_key_auth)], + hibp_client: Annotated[AsyncHTTPHandler, Depends(get_hibp_client)], ) -> ChangePasswordResponse: """ Change the calling user's own password. @@ -133,7 +139,7 @@ async def change_password( ) validate_password_policy(data.new_password, general_settings) - await validate_password_not_breached(data.new_password, general_settings) + await validate_password_not_breached(data.new_password, general_settings, hibp_client) password_update: Final[prisma_types.LiteLLM_UserTableUpdateInput] = { "password": hash_password(data.new_password), diff --git a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py index b260d240f29..ae30b086c6e 100644 --- a/tests/test_litellm/experimental_mcp_client/test_mcp_client.py +++ b/tests/test_litellm/experimental_mcp_client/test_mcp_client.py @@ -1890,8 +1890,9 @@ async def test_transport_completion_and_normal_messages(transport: MCPTransport, from litellm.proxy._experimental.mcp_server.rest_endpoints import _connection_error_message logging_callback: Final = AsyncMock() + read_timeout: Final = 0.2 if mode == "silent" else 30 client: Final = MCPClient( - server_url="https://example.com/sse", transport_type=transport, timeout=0.2, logging_callback=logging_callback + server_url="https://example.com/sse", transport_type=transport, timeout=read_timeout, logging_callback=logging_callback ) async def operation(session: ClientSession) -> CallToolResult: @@ -1909,7 +1910,7 @@ async def test_transport_completion_and_normal_messages(transport: MCPTransport, with pytest.raises(MCPError) as caught: await asyncio.wait_for(pending, timeout=3) if mode == "closed": - assert "connection was closed" in _connection_error_message(caught.value, client.server_url, 0.2) + assert "connection was closed" in _connection_error_message(caught.value, client.server_url, read_timeout) else: assert isinstance(as_mcp_read_timeout(caught.value), TimeoutError) diff --git a/tests/test_litellm/proxy/client/cli/autoroute/test_wizard.py b/tests/test_litellm/proxy/client/cli/autoroute/test_wizard.py index fc6de53cb9e..c0e5377b170 100644 --- a/tests/test_litellm/proxy/client/cli/autoroute/test_wizard.py +++ b/tests/test_litellm/proxy/client/cli/autoroute/test_wizard.py @@ -1,4 +1,6 @@ import asyncio +import contextlib +import contextvars from typing import Any, Dict, List, Optional, Tuple from unittest.mock import patch @@ -288,9 +290,12 @@ def _highlighted_choice(session: AppSession) -> Optional[str]: if session.app is None: return None controls = [c for c in session.app.layout.find_all_controls() if isinstance(c, InquirerPyFuzzyControl)] - if not controls or controls[0].choice_count == 0: + if not controls: + return None + try: + return controls[0].selection["name"] + except IndexError: return None - return controls[0].selection["name"] async def _wait_until_highlighted(session: AppSession, name: str) -> None: @@ -309,21 +314,35 @@ def _drive_fuzzy_pick( ) -> List[str]: """Drives the real InquirerPy fuzzy prompt through prompt_toolkit's own test input/output, exercising the actual widget (filtering, tab-to-toggle, enter-to-confirm) rather than mocking - it away. asyncio.to_thread propagates the create_app_session context into the worker thread - running _fuzzy_pick's synchronous .execute() call. Each key event names the choice the widget - must highlight before the next key is sent (None sends the next key immediately).""" + it away. The worker thread running _fuzzy_pick's synchronous .execute() call inherits the + create_app_session context. Each key event names the choice the widget must highlight before + the next key is sent (None sends the next key immediately). The widget swaps its filtered list + before it clamps the highlight index on the next redraw, so the poller only reads a name once + the index is in range. If driving the widget fails, ctrl-c ends the prompt so the worker thread + exits and the failure surfaces instead of hanging the event loop shutdown.""" async def _run() -> List[str]: with create_pipe_input() as pipe_input: with create_app_session(input=pipe_input, output=DummyOutput()) as session: - task = asyncio.ensure_future( - asyncio.to_thread(wizard_module._fuzzy_pick, models, prompt_label, multiselect) + prompt = asyncio.get_running_loop().run_in_executor( + None, + contextvars.copy_context().run, + wizard_module._fuzzy_pick, + models, + prompt_label, + multiselect, ) - for text, highlighted in key_events: - pipe_input.send_text(text) - if highlighted is not None: - await _wait_until_highlighted(session, highlighted) - return await task + try: + for text, highlighted in key_events: + pipe_input.send_text(text) + if highlighted is not None: + await _wait_until_highlighted(session, highlighted) + except BaseException: + pipe_input.send_text("\x03") + with contextlib.suppress(BaseException): + await prompt + raise + return await prompt return asyncio.run(_run()) diff --git a/tests/test_litellm/proxy/management_endpoints/test_password_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_password_endpoints.py index 984c95321b5..adff4eda47c 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_password_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_password_endpoints.py @@ -1,17 +1,19 @@ """ Tests for POST /user/password/change (litellm/proxy/management_endpoints/password_endpoints.py). -HIBP traffic is intercepted with respx; no test here touches the network. +HIBP is served by an AsyncHTTPHandler wrapping an httpx.MockTransport that is +injected straight into change_password; no test here touches the network. """ import hashlib +from typing import Final from unittest.mock import AsyncMock, MagicMock, patch import httpx import pytest -import respx from fastapi import HTTPException +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.proxy._types import UI_TEAM_ID, LitellmTableNames, ProxyErrorTypes, ProxyException, UserAPIKeyAuth from litellm.proxy.auth.login_utils import PASSWORD_SESSION_METADATA from litellm.proxy.management_endpoints.password_endpoints import change_password @@ -49,15 +51,29 @@ def _virtual_key_caller() -> UserAPIKeyAuth: return UserAPIKeyAuth(user_id="user-123", team_id="team-abc", metadata=dict(PASSWORD_SESSION_METADATA)) -def _hibp_url_for(password: str) -> str: - sha1 = hashlib.sha1(password.encode(), usedforsecurity=False).hexdigest().upper() - return f"https://api.pwnedpasswords.com/range/{sha1[:5]}" - - def _hibp_suffix_for(password: str) -> str: return hashlib.sha1(password.encode(), usedforsecurity=False).hexdigest().upper()[5:] +def _hibp_client_returning(body: str) -> AsyncHTTPHandler: + return AsyncHTTPHandler(transport=httpx.MockTransport(lambda request: httpx.Response(200, text=body))) + + +def _hibp_client_never_called() -> AsyncHTTPHandler: + def handler(request: httpx.Request) -> httpx.Response: + raise AssertionError(f"unexpected HIBP call to {request.url}") + + return AsyncHTTPHandler(transport=httpx.MockTransport(handler)) + + +def _hibp_client_recording(calls: list[httpx.Request]) -> AsyncHTTPHandler: + def handler(request: httpx.Request) -> httpx.Response: + calls.append(request) + return httpx.Response(200, text="") + + return AsyncHTTPHandler(transport=httpx.MockTransport(handler)) + + @pytest.mark.asyncio async def test_change_password_success_writes_new_scrypt_hash(): from litellm.proxy._types import ChangePasswordRequest @@ -75,6 +91,7 @@ async def test_change_password_success_writes_new_scrypt_hash(): response = await change_password( data=ChangePasswordRequest(current_password=CURRENT_PASSWORD, new_password=NEW_PASSWORD), user_api_key_dict=_caller(), + hibp_client=_hibp_client_never_called(), ) assert response.user_id == "user-123" @@ -107,6 +124,7 @@ async def test_change_password_rejects_wrong_current_password(): await change_password( data=ChangePasswordRequest(current_password="not-the-password", new_password=NEW_PASSWORD), user_api_key_dict=_caller(), + hibp_client=_hibp_client_never_called(), ) assert exc_info.value.status_code == 400 @@ -132,6 +150,7 @@ async def test_change_password_rejects_unchanged_password(): await change_password( data=ChangePasswordRequest(current_password=CURRENT_PASSWORD, new_password=CURRENT_PASSWORD), user_api_key_dict=_caller(), + hibp_client=_hibp_client_never_called(), ) assert exc_info.value.status_code == 400 @@ -167,6 +186,7 @@ async def test_change_password_rejects_non_password_login_session(caller: UserAP await change_password( data=ChangePasswordRequest(current_password=CURRENT_PASSWORD, new_password=NEW_PASSWORD), user_api_key_dict=caller, + hibp_client=_hibp_client_never_called(), ) assert exc_info.value.status_code == 403 @@ -193,6 +213,7 @@ async def test_change_password_rejects_session_without_user(): await change_password( data=ChangePasswordRequest(current_password=CURRENT_PASSWORD, new_password=NEW_PASSWORD), user_api_key_dict=_caller(user_id=None), + hibp_client=_hibp_client_never_called(), ) assert exc_info.value.status_code == 400 @@ -219,6 +240,7 @@ async def test_change_password_rejects_account_without_password(): await change_password( data=ChangePasswordRequest(current_password=CURRENT_PASSWORD, new_password=NEW_PASSWORD), user_api_key_dict=_caller(), + hibp_client=_hibp_client_never_called(), ) assert exc_info.value.status_code == 400 @@ -244,6 +266,7 @@ async def test_change_password_enforces_min_length(): await change_password( data=ChangePasswordRequest(current_password=CURRENT_PASSWORD, new_password="Short1!"), user_api_key_dict=_caller(), + hibp_client=_hibp_client_never_called(), ) assert exc_info.value.code == "400" @@ -254,15 +277,11 @@ async def test_change_password_enforces_min_length(): @pytest.mark.asyncio -@respx.mock async def test_change_password_rejects_breached_password(): """With the default policy, the new password is screened against HIBP.""" from litellm.proxy._types import ChangePasswordRequest breached_password = "Password123!" - respx.get(_hibp_url_for(breached_password)).mock( - return_value=httpx.Response(200, text=f"{_hibp_suffix_for(breached_password)}:1") - ) prisma = _make_prisma(_make_user_row(hash_password(CURRENT_PASSWORD))) with ( @@ -277,6 +296,7 @@ async def test_change_password_rejects_breached_password(): await change_password( data=ChangePasswordRequest(current_password=CURRENT_PASSWORD, new_password=breached_password), user_api_key_dict=_caller(), + hibp_client=_hibp_client_returning(f"{_hibp_suffix_for(breached_password)}:1"), ) assert exc_info.value.code == "400" @@ -287,16 +307,14 @@ async def test_change_password_rejects_breached_password(): @pytest.mark.asyncio -@respx.mock async def test_change_password_verifies_current_password_before_hibp_lookup(): """A caller who fails current-password verification must not trigger any - HIBP traffic. The HIBP check fails open on errors, so an unmocked lookup - could not prove ordering; instead the route is registered and asserted - uncalled.""" + HIBP traffic: the injected client records each request it serves and the + test asserts none were made.""" from litellm.proxy._types import ChangePasswordRequest - hibp_route = respx.get(_hibp_url_for(NEW_PASSWORD)).mock(return_value=httpx.Response(200, text="")) prisma = _make_prisma(_make_user_row(hash_password(CURRENT_PASSWORD))) + hibp_calls: Final[list[httpx.Request]] = [] with ( patch( # test-quality-ok: change_password reads proxy_server module globals; no injection seam @@ -310,11 +328,12 @@ async def test_change_password_verifies_current_password_before_hibp_lookup(): await change_password( data=ChangePasswordRequest(current_password="not-the-password", new_password=NEW_PASSWORD), user_api_key_dict=_caller(), + hibp_client=_hibp_client_recording(hibp_calls), ) assert exc_info.value.status_code == 400 assert "Current password is incorrect" in exc_info.value.detail["error"] - assert not hibp_route.called + assert hibp_calls == [] prisma.db.litellm_usertable.update.assert_not_called() @@ -341,6 +360,7 @@ async def test_change_password_success_emits_redacted_audit_log(): await change_password( data=ChangePasswordRequest(current_password=CURRENT_PASSWORD, new_password=NEW_PASSWORD), user_api_key_dict=_caller(), + hibp_client=_hibp_client_never_called(), ) audit_mock.assert_awaited_once() @@ -375,6 +395,7 @@ async def test_change_password_failure_emits_no_audit_log(): await change_password( data=ChangePasswordRequest(current_password="not-the-password", new_password=NEW_PASSWORD), user_api_key_dict=_caller(), + hibp_client=_hibp_client_never_called(), ) audit_mock.assert_not_awaited() @@ -411,6 +432,7 @@ async def test_change_password_revokes_other_sessions_keeping_callers(): await change_password( data=ChangePasswordRequest(current_password=CURRENT_PASSWORD, new_password=NEW_PASSWORD), user_api_key_dict=caller, + hibp_client=_hibp_client_never_called(), ) revoke_mock.assert_awaited_once() @@ -442,6 +464,7 @@ async def test_change_password_failure_revokes_no_sessions(): await change_password( data=ChangePasswordRequest(current_password="not-the-password", new_password=NEW_PASSWORD), user_api_key_dict=_caller(), + hibp_client=_hibp_client_never_called(), ) revoke_mock.assert_not_awaited() @@ -463,6 +486,7 @@ async def test_change_password_requires_db(): await change_password( data=ChangePasswordRequest(current_password=CURRENT_PASSWORD, new_password=NEW_PASSWORD), user_api_key_dict=_caller(), + hibp_client=_hibp_client_never_called(), ) assert exc_info.value.status_code == 500 From 21530d887b465393e80fad09dd01588f8b44c67b Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 09:28:09 -0700 Subject: [PATCH 07/10] feat(gemini): add gemini-3.8-flash-tts and gemini-3.8-flash-lite-tts prices (#42752) * feat(gemini): add gemini-3.8-flash-tts and gemini-3.8-flash-lite-tts prices Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(gemini): bill tiered TTS output through output_cost_per_token tiers Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- ...odel_prices_and_context_window_backup.json | 48 +++++++++++++++++++ model_prices_and_context_window.json | 48 +++++++++++++++++++ tests/test_litellm/test_utils.py | 2 + 3 files changed, 98 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index 345daa5c8a3..cc3bf192908 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -55199,6 +55199,54 @@ "supports_response_schema": false, "supports_web_search": false }, + "gemini/gemini-3.8-flash-tts": { + "cache_read_input_token_cost": 1.25e-07, + "cache_read_input_token_cost_batches": 6.25e-08, + "cache_read_input_token_cost_flex": 2.5e-08, + "cache_read_input_token_cost_priority": 2.25e-07, + "input_cost_per_token": 5e-07, + "input_cost_per_token_batches": 2.5e-07, + "input_cost_per_token_flex": 2.5e-07, + "input_cost_per_token_priority": 9e-07, + "litellm_provider": "gemini", + "max_input_tokens": 8192, + "max_output_tokens": 16384, + "max_tokens": 16384, + "mode": "audio_speech", + "output_cost_per_audio_token": 9e-06, + "output_cost_per_token": 9e-06, + "output_cost_per_token_batches": 4.5e-06, + "output_cost_per_token_flex": 4.5e-06, + "output_cost_per_token_priority": 1.62e-05, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/audio/speech" + ] + }, + "gemini/gemini-3.8-flash-lite-tts": { + "cache_read_input_token_cost": 1.25e-07, + "cache_read_input_token_cost_batches": 6.25e-08, + "cache_read_input_token_cost_flex": 2.5e-08, + "cache_read_input_token_cost_priority": 2.25e-07, + "input_cost_per_token": 5e-07, + "input_cost_per_token_batches": 2.5e-07, + "input_cost_per_token_flex": 2.5e-07, + "input_cost_per_token_priority": 9e-07, + "litellm_provider": "gemini", + "max_input_tokens": 8192, + "max_output_tokens": 16384, + "max_tokens": 16384, + "mode": "audio_speech", + "output_cost_per_audio_token": 6e-06, + "output_cost_per_token": 6e-06, + "output_cost_per_token_batches": 3e-06, + "output_cost_per_token_flex": 3e-06, + "output_cost_per_token_priority": 1.08e-05, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/audio/speech" + ] + }, "gemini-2.5-flash-preview-tts": { "input_cost_per_token": 5e-07, "input_cost_per_token_batches": 2.5e-07, diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index 345daa5c8a3..cc3bf192908 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -55199,6 +55199,54 @@ "supports_response_schema": false, "supports_web_search": false }, + "gemini/gemini-3.8-flash-tts": { + "cache_read_input_token_cost": 1.25e-07, + "cache_read_input_token_cost_batches": 6.25e-08, + "cache_read_input_token_cost_flex": 2.5e-08, + "cache_read_input_token_cost_priority": 2.25e-07, + "input_cost_per_token": 5e-07, + "input_cost_per_token_batches": 2.5e-07, + "input_cost_per_token_flex": 2.5e-07, + "input_cost_per_token_priority": 9e-07, + "litellm_provider": "gemini", + "max_input_tokens": 8192, + "max_output_tokens": 16384, + "max_tokens": 16384, + "mode": "audio_speech", + "output_cost_per_audio_token": 9e-06, + "output_cost_per_token": 9e-06, + "output_cost_per_token_batches": 4.5e-06, + "output_cost_per_token_flex": 4.5e-06, + "output_cost_per_token_priority": 1.62e-05, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/audio/speech" + ] + }, + "gemini/gemini-3.8-flash-lite-tts": { + "cache_read_input_token_cost": 1.25e-07, + "cache_read_input_token_cost_batches": 6.25e-08, + "cache_read_input_token_cost_flex": 2.5e-08, + "cache_read_input_token_cost_priority": 2.25e-07, + "input_cost_per_token": 5e-07, + "input_cost_per_token_batches": 2.5e-07, + "input_cost_per_token_flex": 2.5e-07, + "input_cost_per_token_priority": 9e-07, + "litellm_provider": "gemini", + "max_input_tokens": 8192, + "max_output_tokens": 16384, + "max_tokens": 16384, + "mode": "audio_speech", + "output_cost_per_audio_token": 6e-06, + "output_cost_per_token": 6e-06, + "output_cost_per_token_batches": 3e-06, + "output_cost_per_token_flex": 3e-06, + "output_cost_per_token_priority": 1.08e-05, + "source": "https://ai.google.dev/gemini-api/docs/pricing", + "supported_endpoints": [ + "/v1/audio/speech" + ] + }, "gemini-2.5-flash-preview-tts": { "input_cost_per_token": 5e-07, "input_cost_per_token_batches": 2.5e-07, diff --git a/tests/test_litellm/test_utils.py b/tests/test_litellm/test_utils.py index ce280cc3513..2a8b31c12cc 100644 --- a/tests/test_litellm/test_utils.py +++ b/tests/test_litellm/test_utils.py @@ -6142,6 +6142,8 @@ def test_get_model_info_gemini(monkeypatch): and "veo" not in model and "lyria" not in model and "robotics" not in model + and "3.8-flash-tts" not in model + and "3.8-flash-lite-tts" not in model ): assert info.get("tpm") is not None, f"{model} does not have tpm" assert info.get("rpm") is not None, f"{model} does not have rpm" From a286ebf42eb1e93aae34b985ed1ef1a5edbbe1b0 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 09:51:06 -0700 Subject: [PATCH 08/10] test(integration): regression tests for July provider translation, routing and streaming bugs (#42693) * test(integration): optional Anthropic tool properties stay optional on the OpenAI Responses wire (Pylon #6619) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): Bedrock InvokeModel count-suffixed cache usage fields are reported and charged (Pylon #6708) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): anthropic messages honors the deployment request timeout (Pylon #6505) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): drop client_metadata before the Bedrock Converse body reaches the provider (Pylon #6645) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): repeat Bedrock requests under one session name assume the role once (Pylon #6681) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): messages stream keeps include_usage off the Responses wire with always_include_stream_usage (Pylon #6466) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): clamp sub-16 max_tokens to the Responses API floor instead of 400 (Pylon #6539) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): forwarded client x- headers reach the provider on /v1/responses (Pylon #6565) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): Anthropic messages stop_sequences reach OpenAI-compatible providers as stop (Pylon #6536) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): Codex namespace tools reach a chat upstream flattened and round-trip through /v1/responses (Pylon #6409) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): keep Claude 4.6 legacy thinking budget_tokens on /v1/messages (Pylon #6727) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): nvidia nim ranking keeps image passages and applies top_n without sending top_k (Pylon #6401) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): tpm-only model rejects priority traffic once recorded tokens reach the model tpm (Pylon #6344) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): reasoning-only chunks open an Anthropic thinking block at index zero on /v1/messages streams (Pylon #6337) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): file content streams to the client before the upstream finishes sending (Pylon #6315) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): agent whose card lives only at agentCard/v1.0 is reached with bearer auth (Pylon #6249) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): vertex batch create returns a batch when outputInfo is null (Pylon #6374) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): fireworks session id is sent as x-session-affinity and cached tokens land in spend log metadata (Pylon #6220) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): advisor sub-call failure does not cool down the executor deployment (Pylon #6212) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): Gemini /v1/messages cache_control creates cachedContent with Anthropic ttl (Pylon #6221) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): bedrock_mantle max_output_tokens below 16 is clamped before reaching Mantle (Pylon #6262) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): missing thinking signature 400 on /v1/messages retries without thinking blocks (Pylon #6222) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): format Gemini messages cache_control wire test (Pylon #6221) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): rebuilt shared aiohttp session keeps the configured keepalive timeout (Pylon #6387) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): sagemaker_chat signs the inference component header and sends hf_model_name as the body model (Pylon #6187) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): Bedrock Converse DeepSeek drops Anthropic thinking and sends V3 reasoning_effort raw (Pylon #6149) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): concurrent team model TPM requests are reserved before the provider call (Pylon #6075) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): /v1/messages honors the configured timeout against a stalled upstream (Pylon #6025) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): Codex additional_tools input items reach Bedrock Mantle as top-level tools (Pylon #6012) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): advisor api_base without api_key never sends the proxy Anthropic key to the caller host (Pylon #6226) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): rerank responses carry call id, latency and cost headers (Pylon #5981) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): parse the outbound Anthropic body with the typed JSON adapter (Pylon #6025) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): Bedrock Knowledge Base search forwards userContext to the Retrieve body (Pylon #5991) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): Marengo 3.0 text embeddings reach Bedrock nested under inputType (Pylon #5949) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): sub-16 max_tokens over a responses deployment reaches OpenAI as 16 (Pylon #6008) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): midturn system correction reaches the OpenAI Responses wire via /v1/messages (Pylon #6449) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): concurrent requests over a key tpm limit are rejected before reaching the provider (Pylon #5737) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): chat to responses bridge keeps deployment AWS credentials for Bedrock Mantle SigV4 (Pylon #5870) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): vertex gemini stream split across many fragments completes without stalling the proxy (Pylon #5838) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): bedrock mantle /v1/messages stream keeps stream true and relays SSE events (Pylon #5596) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): large chat payloads are released from worker memory after the request ends (Pylon #5920) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): streaming success logs v3 rate limit remaining values for callbacks (Pylon #5767) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): a database created search tool backs Anthropic web search interception (Pylon #5669) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): register july provider regression contracts Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): make july provider regression tests deterministic under cache and worker sharing Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): drop order-fragile worker memory probe pending a real retention regression check Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): apply ruff import sorting and formatting Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): drop stale contract entry and pass question to advisor executor Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): use tiktoken-backed executor model in advisor tests Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../compatibility/test_a2a_wire_versions.py | 101 ++++++++++- .../providers/test_anthropic_advisor_wire.py | 124 ++++++++++++-- ...t_anthropic_legacy_thinking_budget_wire.py | 77 +++++++++ ..._anthropic_messages_fireworks_stop_wire.py | 65 +++++++ ...t_anthropic_messages_openai_bridge_wire.py | 82 +++++++++ ...st_anthropic_messages_openai_tools_wire.py | 92 ++++++++++ .../test_anthropic_messages_timeout_wire.py | 58 +++++++ ...anthropic_thinking_signature_retry_wire.py | 95 +++++++++++ .../providers/test_anthropic_wire.py | 161 +++++++++++++++--- ...t_bedrock_converse_client_metadata_wire.py | 43 +++++ .../test_bedrock_deepseek_reasoning_wire.py | 93 ++++++++++ .../test_bedrock_invoke_cache_usage_wire.py | 86 ++++++++++ ...edrock_knowledge_base_user_context_wire.py | 81 +++++++++ .../test_bedrock_mantle_codex_input_wire.py | 53 ++++++ .../test_bedrock_mantle_responses_wire.py | 40 +++++ .../providers/test_bedrock_mantle_wire.py | 141 +++++++++++++++ .../test_bedrock_marengo_embed_3_wire.py | 35 ++++ .../test_bedrock_role_configuration.py | 154 +++++++++++++++-- .../test_bedrock_thinking_tokens_wire.py | 15 +- ...test_fireworks_ai_session_affinity_wire.py | 69 ++++++++ ...test_gemini_messages_cache_control_wire.py | 86 ++++++++++ .../providers/test_nvidia_nim_ranking_wire.py | 47 +++++ .../test_rerank_latency_headers_wire.py | 50 ++++++ .../test_responses_bridge_incomplete.py | 126 ++++++++++++++ .../test_responses_bridge_namespace_tools.py | 159 +++++++++++++++++ .../test_responses_bridge_stream_options.py | 98 +++++++++++ ...responses_client_header_forwarding_wire.py | 87 ++++++++++ .../providers/test_sagemaker_chat_wire.py | 90 ++++++++++ .../test_vertex_batch_output_info_wire.py | 124 ++++++++++++++ ...st_vertex_gemini_fragmented_stream_wire.py | 138 +++++++++++++++ .../test_websearch_interception_wire.py | 110 ++++++++++++ .../routing/test_advisor_failure_cooldown.py | 101 +++++++++++ .../routing/test_key_tpm_reservation.py | 59 +++++++ .../test_priority_model_tpm_enforcement.py | 116 +++++++++++++ .../test_priority_rate_limit_headers.py | 109 +++++++++++- .../routing/test_team_model_tpm_limit.py | 89 ++++++++++ .../sdk/test_aiohttp_session_rebuild_wire.py | 109 ++++++++++++ .../streaming/test_file_content_streaming.py | 48 ++++++ .../streaming/test_stream_contracts.py | 85 ++++++++- 39 files changed, 3434 insertions(+), 62 deletions(-) create mode 100644 tests/integration/providers/test_anthropic_legacy_thinking_budget_wire.py create mode 100644 tests/integration/providers/test_anthropic_messages_fireworks_stop_wire.py create mode 100644 tests/integration/providers/test_anthropic_messages_openai_bridge_wire.py create mode 100644 tests/integration/providers/test_anthropic_messages_openai_tools_wire.py create mode 100644 tests/integration/providers/test_anthropic_messages_timeout_wire.py create mode 100644 tests/integration/providers/test_anthropic_thinking_signature_retry_wire.py create mode 100644 tests/integration/providers/test_bedrock_converse_client_metadata_wire.py create mode 100644 tests/integration/providers/test_bedrock_deepseek_reasoning_wire.py create mode 100644 tests/integration/providers/test_bedrock_invoke_cache_usage_wire.py create mode 100644 tests/integration/providers/test_bedrock_knowledge_base_user_context_wire.py create mode 100644 tests/integration/providers/test_bedrock_marengo_embed_3_wire.py create mode 100644 tests/integration/providers/test_fireworks_ai_session_affinity_wire.py create mode 100644 tests/integration/providers/test_gemini_messages_cache_control_wire.py create mode 100644 tests/integration/providers/test_nvidia_nim_ranking_wire.py create mode 100644 tests/integration/providers/test_rerank_latency_headers_wire.py create mode 100644 tests/integration/providers/test_responses_bridge_namespace_tools.py create mode 100644 tests/integration/providers/test_responses_bridge_stream_options.py create mode 100644 tests/integration/providers/test_responses_client_header_forwarding_wire.py create mode 100644 tests/integration/providers/test_sagemaker_chat_wire.py create mode 100644 tests/integration/providers/test_vertex_batch_output_info_wire.py create mode 100644 tests/integration/providers/test_vertex_gemini_fragmented_stream_wire.py create mode 100644 tests/integration/routing/test_advisor_failure_cooldown.py create mode 100644 tests/integration/routing/test_key_tpm_reservation.py create mode 100644 tests/integration/routing/test_priority_model_tpm_enforcement.py create mode 100644 tests/integration/routing/test_team_model_tpm_limit.py create mode 100644 tests/integration/sdk/test_aiohttp_session_rebuild_wire.py create mode 100644 tests/integration/streaming/test_file_content_streaming.py diff --git a/tests/integration/compatibility/test_a2a_wire_versions.py b/tests/integration/compatibility/test_a2a_wire_versions.py index 7a828ba2487..2823911ace3 100644 --- a/tests/integration/compatibility/test_a2a_wire_versions.py +++ b/tests/integration/compatibility/test_a2a_wire_versions.py @@ -3,7 +3,6 @@ import uuid from typing import Final import pytest - from integration._support.client import Gateway from integration._support.database import read_rows from integration._support.wire import Reply, Request, wire_server @@ -115,3 +114,103 @@ def test_a2a_versions_and_legacy_casing_preserve_real_wire_and_response(gateway: actual: Final = wire.drain() assert len(tuple(item for item in actual if item.method == "POST")) == 1 assert any(item.method == "GET" for item in actual) + + +@pytest.mark.covers("compatibility.a2a.versioned_card_path_agent_is_reached_with_bearer_and_blocking_send") +def test_agent_serving_its_card_only_at_versioned_path_is_reached_with_bearer_and_answers(gateway: Gateway) -> None: + marker: Final = "foundry" + uuid.uuid4().hex + bearer: Final = "Bearer synthetic-entra-" + marker + + def upstream(request: Request) -> Reply: + assert request.headers.get("authorization") == bearer, request.headers + if request.method == "GET": + if request.target != "/agentCard/v1.0": + return Reply(status=404, body=json.dumps({"error": "not found"}).encode()) + return Reply( + body=json.dumps( + { + "protocolVersion": "0.3", + "name": marker, + "description": "Synthetic prompt agent", + "version": "1.0.0", + "url": wire.url + "/", + "capabilities": {"streaming": False}, + "defaultInputModes": ["text"], + "defaultOutputModes": ["text"], + "skills": [], + } + ).encode() + ) + assert request.method == "POST" and request.target == "/", request.target + body: Final = json.loads(request.body) + assert body["jsonrpc"] == "2.0" and body["method"] == "message/send", body + message: Final = body["params"]["message"] + assert message["kind"] == "message" and message["role"] == "user", message + assert message["parts"] == [{"kind": "text", "text": "synthetic ping"}], message + return Reply( + body=json.dumps( + { + "jsonrpc": "2.0", + "id": body["id"], + "result": { + "kind": "message", + "role": "agent", + "messageId": marker + "-out", + "parts": [{"kind": "text", "text": "synthetic pong"}], + }, + } + ).encode() + ) + + with wire_server(upstream) as wire, gateway.scenario() as scenario: + card: Final = { + "protocolVersion": "0.3", + "name": marker, + "description": "Synthetic prompt agent", + "version": "1.0.0", + "url": wire.url + "/", + "capabilities": {"streaming": False}, + "defaultInputModes": ["text"], + "defaultOutputModes": ["text"], + "skills": [], + } + created: Final = gateway.request( + "POST", + "/v1/agents", + {"agent_name": marker, "agent_card_params": card, "static_headers": {"Authorization": bearer}}, + ) + assert created.status_code == 200, created.text + identity: Final = created.json()["agent_id"] + + def cleanup() -> None: + deleted: Final = gateway.request("DELETE", f"/v1/agents/{identity}") + assert deleted.status_code == 200, deleted.text + assert read_rows('SELECT agent_id FROM "LiteLLM_AgentsTable" WHERE agent_id=%s', (identity,)) == [] + + scenario.cleanups.callback(cleanup) + response: Final = gateway.request( + "POST", + f"/a2a/{identity}", + { + "jsonrpc": "2.0", + "id": marker, + "method": "message/send", + "params": { + "message": { + "kind": "message", + "role": "user", + "messageId": marker + "-in", + "parts": [{"kind": "text", "text": "synthetic ping"}], + } + }, + }, + ) + assert response.status_code == 200, response.text + body: Final = response.json() + assert body["jsonrpc"] == "2.0" and body["id"] == marker and "error" not in body, response.text + assert body["result"]["kind"] == "message", response.text + assert body["result"]["messageId"] == marker + "-out", response.text + assert body["result"]["parts"] == [{"kind": "text", "text": "synthetic pong"}], response.text + actual: Final = wire.drain() + assert tuple(item.target for item in actual if item.method == "GET")[-1] == "/agentCard/v1.0", actual + assert tuple(item.target for item in actual if item.method == "POST") == ("/",), actual diff --git a/tests/integration/providers/test_anthropic_advisor_wire.py b/tests/integration/providers/test_anthropic_advisor_wire.py index 77fa27cd2a9..b2f44d8f155 100644 --- a/tests/integration/providers/test_anthropic_advisor_wire.py +++ b/tests/integration/providers/test_anthropic_advisor_wire.py @@ -1,28 +1,35 @@ import json import uuid +from pathlib import Path from typing import Final import pytest +import yaml from integration._support.client import Gateway +from integration._support.process import owned_proxy from integration._support.wire import Reply, Request, wire_server _ADVISOR_KEY: Final = "synthetic-advisor-key" +_PROXY_ANTHROPIC_KEY: Final = "sk-proxy-owned-anthropic-secret" _QUESTION: Final = "which index should this query use" _ADVICE: Final = "use the composite index on (tenant_id, created_at)" _FINAL_ANSWER: Final = "done, the composite index is the right one" -_ADVISOR_CALL_MESSAGE: Final = { - "role": "assistant", - "content": None, - "tool_calls": [ - { - "id": "advisor-call", - "type": "function", - "function": {"name": "advisor", "arguments": json.dumps({"question": _QUESTION})}, - } - ], -} +def _advisor_call_message(question: str) -> dict[str, object]: + return { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "advisor-call", + "type": "function", + "function": {"name": "advisor", "arguments": json.dumps({"question": question})}, + } + ], + } + + _FINAL_MESSAGE: Final = {"role": "assistant", "content": _FINAL_ANSWER} @@ -41,7 +48,7 @@ def _chat_completion(identity: str, message: dict[str, object], finish_reason: s ) -def _executor_reply(body: dict[str, object], identity: str) -> Reply: +def _executor_reply(body: dict[str, object], identity: str, question: str) -> Reply: messages: Final = body["messages"] assert isinstance(messages, list) if any(message.get("role") == "tool" for message in messages): @@ -50,7 +57,7 @@ def _executor_reply(body: dict[str, object], identity: str) -> Reply: tools: Final = body["tools"] assert isinstance(tools, list) assert tools[0]["function"]["name"] == "advisor" - return _chat_completion(identity, _ADVISOR_CALL_MESSAGE, "tool_calls") + return _chat_completion(identity, _advisor_call_message(question), "tool_calls") @pytest.mark.covers("providers.anthropic_messages_advisor.sub_call_uses_the_configured_advisor_deployment") @@ -58,18 +65,20 @@ def test_advisor_sub_call_reaches_the_router_deployment_with_its_key_instead_of_ gateway: Gateway, ) -> None: identity: Final = "advisor-wire-" + uuid.uuid4().hex + migration: Final = "please plan the migration " + identity + question: Final = _QUESTION + " " + identity def respond(request: Request) -> Reply: body: Final = json.loads(request.body) if request.target == "/v1/chat/completions": assert request.headers["authorization"] == "Bearer integration-provider-key" - return _executor_reply(body, identity) + return _executor_reply(body, identity, question) assert request.target == "/v1/messages" assert request.headers["x-api-key"] == _ADVISOR_KEY assert body["model"] == "claude-opus-4-1-20250805" assert body["messages"] == [ - {"role": "user", "content": "please plan the migration"}, - {"role": "user", "content": _QUESTION}, + {"role": "user", "content": migration}, + {"role": "user", "content": question}, ] assert "tools" not in body return Reply( @@ -88,7 +97,7 @@ def test_advisor_sub_call_reaches_the_router_deployment_with_its_key_instead_of_ ) with wire_server(respond) as wire, gateway.scenario() as scenario: - executor: Final = scenario.model(model="hosted_vllm/llama-3.3-70b", api_base=wire.url + "/v1") + executor: Final = scenario.model(model="hosted_vllm/gpt-4o-mini", api_base=wire.url + "/v1") advisor: Final = scenario.model( model="anthropic/claude-opus-4-1-20250805", api_base=wire.url, api_key=_ADVISOR_KEY ) @@ -98,7 +107,7 @@ def test_advisor_sub_call_reaches_the_router_deployment_with_its_key_instead_of_ { "model": executor, "max_tokens": 64, - "messages": [{"role": "user", "content": "please plan the migration"}], + "messages": [{"role": "user", "content": migration}], "tools": [{"type": "advisor_20260301", "name": "advisor", "model": advisor}], }, ) @@ -111,3 +120,82 @@ def test_advisor_sub_call_reaches_the_router_deployment_with_its_key_instead_of_ "/v1/messages", "/v1/chat/completions", ] + + +def _advice_reply(identity: str) -> Reply: + return Reply( + body=json.dumps( + { + "id": f"msg-{identity}", + "type": "message", + "role": "assistant", + "model": "claude-opus-4-1-20250805", + "content": [{"type": "text", "text": _ADVICE}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 12, "output_tokens": 6}, + } + ).encode() + ) + + +@pytest.mark.covers("providers.anthropic_messages_advisor.caller_api_base_without_api_key_never_receives_the_proxy_key") +def test_advisor_api_base_without_api_key_is_rejected_before_the_proxy_anthropic_key_reaches_the_caller_host( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = "advisor-leak-" + uuid.uuid4().hex + question: Final = _QUESTION + " " + identity + + def executor(request: Request) -> Reply: + assert request.target == "/v1/chat/completions", request.target + return _executor_reply(json.loads(request.body), identity, question) + + def caller_host(request: Request) -> Reply: + return _advice_reply(identity) + + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["general_settings"]["allow_client_side_credentials"] = True + path: Final = tmp_path / "client-side-credentials.yaml" + path.write_text(yaml.safe_dump(config)) + with ( + wire_server(executor) as executor_wire, + wire_server(caller_host) as caller_wire, + owned_proxy(gateway, tmp_path, {"ANTHROPIC_API_KEY": _PROXY_ANTHROPIC_KEY}, config=path) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(model="hosted_vllm/gpt-4o-mini", api_base=executor_wire.url + "/v1") + response: Final = candidate.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 64, + "messages": [{"role": "user", "content": "please plan the migration"}], + "tools": [ + { + "type": "advisor_20260301", + "name": "advisor", + "model": "anthropic/claude-opus-4-1-20250805", + "api_base": caller_wire.url, + } + ], + }, + ) + received: Final = caller_wire.drain() + assert [ + (request.target, request.headers.get("x-api-key"), json.loads(request.body)["messages"]) + for request in received + ] == [], response.text + assert response.is_error, response.text + assert response.json() == { + "type": "error", + "error": { + "type": "api_error", + "message": ( + "advisor tool definition sets 'api_base' without 'api_key'. A caller-supplied api_base is only " + "honored alongside a caller-supplied api_key, so the proxy's own credentials are never sent to a " + "caller-chosen destination." + ), + }, + }, response.text + assert executor_wire.drain() == (), response.text diff --git a/tests/integration/providers/test_anthropic_legacy_thinking_budget_wire.py b/tests/integration/providers/test_anthropic_legacy_thinking_budget_wire.py new file mode 100644 index 00000000000..242e5c7ec5a --- /dev/null +++ b/tests/integration/providers/test_anthropic_legacy_thinking_budget_wire.py @@ -0,0 +1,77 @@ +import json +import uuid +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server + +_MODEL: Final = "claude-sonnet-4-6" +_KEY: Final = "synthetic-anthropic-key" +_THINKING: Final = {"type": "enabled", "budget_tokens": 8000} +_TOOL: Final = { + "name": "read_file", + "description": "read a file", + "input_schema": {"type": "object", "properties": {"path": {"type": "string"}}, "required": ["path"]}, +} +_NEXT_CALL: Final = {"type": "tool_use", "id": "call-2", "name": "read_file", "input": {"path": "schema.prisma"}} + + +def _tool_loop_history(identity: str) -> tuple[dict[str, object], ...]: + return ( + {"role": "user", "content": f"open the config for {identity}"}, + { + "role": "assistant", + "content": [{"type": "tool_use", "id": "call-1", "name": "read_file", "input": {"path": "config.yaml"}}], + }, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "call-1", "content": "model_list: []"}]}, + ) + + +def _tool_use_reply(identity: str) -> Reply: + return Reply( + body=json.dumps( + { + "id": f"msg-{identity}", + "type": "message", + "role": "assistant", + "model": _MODEL, + "content": [_NEXT_CALL], + "stop_reason": "tool_use", + "stop_sequence": None, + "usage": {"input_tokens": 40, "output_tokens": 12}, + } + ).encode() + ) + + +@pytest.mark.covers("providers.anthropic_messages.claude_4_6_legacy_thinking_budget_reaches_the_wire_unchanged") +def test_claude_4_6_thinking_budget_tokens_on_messages_is_forwarded_instead_of_rewritten_to_adaptive( + gateway: Gateway, +) -> None: + identity: Final = "legacy-thinking-" + uuid.uuid4().hex + history: Final = _tool_loop_history(identity) + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/messages", request.target + assert request.headers["x-api-key"] == _KEY + body: Final = json.loads(request.body) + assert body["thinking"] == _THINKING, body + assert "output_config" not in body, body + assert body["max_tokens"] == 32768, body + assert body["messages"] == list(history), body + assert body["tools"] == [_TOOL], body + return _tool_use_reply(identity) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"anthropic/{_MODEL}", api_base=wire.url, api_key=_KEY) + response: Final = gateway.request( + "POST", + "/v1/messages", + {"model": model, "max_tokens": 32768, "thinking": _THINKING, "messages": history, "tools": [_TOOL]}, + ) + assert response.status_code == 200, response.text + body: Final = response.json() + assert body["content"] == [_NEXT_CALL], response.text + assert body["stop_reason"] == "tool_use", response.text + assert len(wire.drain()) == 1 diff --git a/tests/integration/providers/test_anthropic_messages_fireworks_stop_wire.py b/tests/integration/providers/test_anthropic_messages_fireworks_stop_wire.py new file mode 100644 index 00000000000..adec5784aa8 --- /dev/null +++ b/tests/integration/providers/test_anthropic_messages_fireworks_stop_wire.py @@ -0,0 +1,65 @@ +import json +import uuid +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_MODEL: Final = "accounts/fireworks/models/glm-5p3" +_API_KEY: Final = "synthetic-fireworks-key" +_STOP: Final = "" +_ANSWER: Final = "allow" +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + + +@pytest.mark.covers( + "providers.anthropic_messages_adapter.stop_sequences_and_disabled_thinking_reach_openai_compatible_provider_as_stop_and_reasoning_effort" +) +def test_messages_stop_sequences_to_fireworks_are_sent_as_stop_not_stop_sequences(gateway: Gateway) -> None: + prompt: Final = "classify this tool call " + uuid.uuid4().hex + + def respond(request: Request) -> Reply: + assert request.method == "POST" + assert request.target == "/chat/completions" + assert request.headers["authorization"] == f"Bearer {_API_KEY}" + body: Final = _JSON_OBJECT.validate_json(request.body) + assert "stop_sequences" not in body, body + assert body["stop"] == [_STOP], body + assert body["reasoning_effort"] == "none", body + assert body["model"] == _MODEL, body + assert body["messages"] == [{"role": "user", "content": prompt}], body + return Reply( + body=json.dumps( + { + "id": "fw-classifier", + "object": "chat.completion", + "created": 1, + "model": _MODEL, + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": _ANSWER}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 9, "completion_tokens": 6, "total_tokens": 15}, + } + ).encode() + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"fireworks_ai/{_MODEL}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 64, + "messages": [{"role": "user", "content": prompt}], + "stop_sequences": [_STOP], + "thinking": {"type": "disabled"}, + }, + ) + assert response.status_code == 200, response.text + payload: Final = _JSON_OBJECT.validate_json(response.content) + assert payload["content"] == [{"type": "text", "text": _ANSWER}], response.text + assert payload["stop_reason"] == "end_turn", response.text + assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/chat/completions")] diff --git a/tests/integration/providers/test_anthropic_messages_openai_bridge_wire.py b/tests/integration/providers/test_anthropic_messages_openai_bridge_wire.py new file mode 100644 index 00000000000..72eba1d89a5 --- /dev/null +++ b/tests/integration/providers/test_anthropic_messages_openai_bridge_wire.py @@ -0,0 +1,82 @@ +import json +import uuid +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "gpt-5.4-mini" +_API_KEY: Final = "synthetic-openai-key" +_CORRECTION: Final = "Stop refactoring the parser and only fix the failing test instead." +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + + +def _responses_reply(identity: str, content: str) -> bytes: + return json.dumps( + { + "id": f"resp_{identity}", + "object": "response", + "created_at": 1789788253, + "status": "completed", + "model": _BACKEND, + "output": [ + { + "type": "message", + "id": f"msg_{identity}", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": content, "annotations": []}], + } + ], + "usage": {"input_tokens": 41, "output_tokens": 5, "total_tokens": 46}, + } + ).encode() + + +@pytest.mark.covers("providers.anthropic_messages_openai_bridge.midturn_system_correction_reaches_the_wire") +def test_midturn_system_correction_is_forwarded_to_openai_responses(gateway: Gateway) -> None: + identity: Final = f"openai-midturn-system-{uuid.uuid4().hex}" + + def respond(request: Request) -> Reply: + assert request.method == "POST" + assert request.target == "/responses" + assert request.headers["authorization"] == f"Bearer {_API_KEY}" + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["model"] == _BACKEND + assert body["instructions"] == "You are a coding agent." + assert body["input"] == [ + {"type": "message", "role": "user", "content": [{"type": "input_text", "text": "Fix the failing test."}]}, + { + "type": "message", + "role": "assistant", + "content": [{"type": "output_text", "text": "I will start by refactoring the parser."}], + }, + {"type": "message", "role": "system", "content": [{"type": "input_text", "text": _CORRECTION}]}, + {"type": "message", "role": "user", "content": [{"type": "input_text", "text": "Continue."}]}, + ], body + return Reply(body=_responses_reply(identity, "Understood.")) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"openai/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 64, + "system": "You are a coding agent.", + "messages": [ + {"role": "user", "content": "Fix the failing test."}, + {"role": "assistant", "content": "I will start by refactoring the parser."}, + {"role": "system", "content": _CORRECTION}, + {"role": "user", "content": "Continue."}, + ], + }, + ) + assert response.status_code == 200, response.text + payload: Final = _JSON_OBJECT.validate_json(response.content) + assert payload["content"] == [{"type": "text", "text": "Understood."}], response.text + assert payload["stop_reason"] == "end_turn", response.text + assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/responses")] diff --git a/tests/integration/providers/test_anthropic_messages_openai_tools_wire.py b/tests/integration/providers/test_anthropic_messages_openai_tools_wire.py new file mode 100644 index 00000000000..605fa45e17b --- /dev/null +++ b/tests/integration/providers/test_anthropic_messages_openai_tools_wire.py @@ -0,0 +1,92 @@ +import json +import uuid +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server + +_BACKEND: Final = "gpt-5.4-mini" +_API_KEY: Final = "synthetic-openai-key" +_TOOL_SCHEMA: Final = { + "type": "object", + "properties": { + "city": {"type": "string"}, + "unit": {"type": "string", "enum": ["celsius", "fahrenheit"]}, + "include_forecast": {"type": "boolean"}, + }, + "required": ["city"], +} + + +@pytest.mark.covers("providers.anthropic_messages_bridge.optional_tool_properties_stay_optional_on_the_wire") +def test_messages_tool_with_optional_properties_reaches_openai_responses_non_strict(gateway: Gateway) -> None: + identity: Final = f"messages-optional-tool-{uuid.uuid4().hex}" + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/responses", request.target + assert request.headers["authorization"] == f"Bearer {_API_KEY}" + body: Final = json.loads(request.body) + assert body["model"] == _BACKEND, body + assert body["tools"] == [ + { + "type": "function", + "name": "get_weather", + "strict": False, + "description": "Current weather for a city", + "parameters": _TOOL_SCHEMA, + } + ], body["tools"] + return Reply( + body=json.dumps( + { + "id": f"resp_{identity}", + "object": "response", + "created_at": 1789788253, + "status": "completed", + "model": _BACKEND, + "output": [ + { + "type": "function_call", + "id": f"fc_{identity}", + "call_id": f"call_{identity}", + "name": "get_weather", + "arguments": json.dumps({"city": "Paris"}), + "status": "completed", + } + ], + "usage": {"input_tokens": 30, "output_tokens": 9, "total_tokens": 39}, + } + ).encode() + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"openai/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 64, + "messages": [{"role": "user", "content": "What is the weather in Paris?"}], + "tools": [ + { + "name": "get_weather", + "description": "Current weather for a city", + "input_schema": _TOOL_SCHEMA, + } + ], + }, + ) + assert response.status_code == 200, response.text + assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/responses")] + body: Final = response.json() + assert body["stop_reason"] == "tool_use", response.text + assert body["content"] == [ + { + "type": "tool_use", + "id": f"call_{identity}", + "name": "get_weather", + "input": {"city": "Paris"}, + } + ], response.text diff --git a/tests/integration/providers/test_anthropic_messages_timeout_wire.py b/tests/integration/providers/test_anthropic_messages_timeout_wire.py new file mode 100644 index 00000000000..76c29ad7763 --- /dev/null +++ b/tests/integration/providers/test_anthropic_messages_timeout_wire.py @@ -0,0 +1,58 @@ +import json +import time +import uuid +from typing import Final + +import pytest +from integration._support.client import JSON_OBJECT, Gateway +from integration._support.wire import Reply, Request, wire_server + +_UPSTREAM_STALL_SECONDS: Final = 4.0 +_CONFIGURED_TIMEOUT_SECONDS: Final = 1.0 + + +@pytest.mark.covers("providers.anthropic_messages.configured_timeout_aborts_stalled_upstream") +def test_messages_endpoint_honors_configured_timeout_against_stalled_upstream(gateway: Gateway) -> None: + prompt: Final = "stall-" + uuid.uuid4().hex + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/messages" + assert request.headers["x-api-key"] == "synthetic-anthropic-key" + body: Final = JSON_OBJECT.validate_json(request.body) + assert body["model"] == "claude-sonnet-4-5-20250929" + assert body["messages"] == [{"role": "user", "content": prompt}] + assert body["max_tokens"] == 16 + assert "timeout" not in body + time.sleep(_UPSTREAM_STALL_SECONDS) + return Reply( + body=json.dumps( + { + "id": "msg_stalled", + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5-20250929", + "content": [{"type": "text", "text": "too late"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 1, "output_tokens": 2}, + } + ).encode() + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", + api_base=wire.url, + api_key="synthetic-anthropic-key", + timeout=_CONFIGURED_TIMEOUT_SECONDS, + ) + started: Final = time.monotonic() + response: Final = gateway.request( + "POST", + "/v1/messages", + {"model": model, "max_tokens": 16, "messages": [{"role": "user", "content": prompt}]}, + ) + elapsed: Final = time.monotonic() - started + assert response.status_code == 408, response.text + assert elapsed < _UPSTREAM_STALL_SECONDS, f"timed out only after {elapsed:.2f}s: {response.text}" + assert len(wire.drain()) == 1 diff --git a/tests/integration/providers/test_anthropic_thinking_signature_retry_wire.py b/tests/integration/providers/test_anthropic_thinking_signature_retry_wire.py new file mode 100644 index 00000000000..e414d8f0d11 --- /dev/null +++ b/tests/integration/providers/test_anthropic_thinking_signature_retry_wire.py @@ -0,0 +1,95 @@ +import json +import uuid +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server + +MODEL: Final = "claude-sonnet-4-5-20250929" +KEY: Final = "synthetic-anthropic-key" +SIGNATURE_ERROR: Final = json.dumps( + { + "type": "error", + "error": { + "type": "invalid_request_error", + "message": "messages.2.content.0.thinking.signature.str: Input should be a valid string", + }, + } +).encode() +TOOLS: Final = ({"name": "lookup", "input_schema": {"type": "object", "properties": {"key": {"type": "string"}}}},) + + +def _history_with_unsigned_thinking(identity: str) -> tuple[dict[str, object], ...]: + return ( + {"role": "user", "content": [{"type": "text", "text": f"first question {identity}"}]}, + {"role": "assistant", "content": [{"type": "text", "text": "first answer"}]}, + {"role": "user", "content": [{"type": "text", "text": "second question"}]}, + { + "role": "assistant", + "content": [ + {"type": "thinking", "thinking": "replayed from another provider", "signature": None}, + {"type": "tool_use", "id": "call-1", "name": "lookup", "input": {"key": "value"}}, + ], + }, + {"role": "user", "content": [{"type": "tool_result", "tool_use_id": "call-1", "content": "found"}]}, + ) + + +@pytest.mark.covers("providers.anthropic_messages.missing_thinking_signature_400_retries_without_thinking_blocks") +def test_missing_thinking_signature_400_retries_once_without_thinking_blocks_and_returns_200( + gateway: Gateway, +) -> None: + identity: Final = "thinking-signature-" + uuid.uuid4().hex + history: Final = _history_with_unsigned_thinking(identity) + tool_use_only_turn: Final = { + "role": "assistant", + "content": [{"type": "tool_use", "id": "call-1", "name": "lookup", "input": {"key": "value"}}], + } + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/messages" + assert request.headers["x-api-key"] == KEY + body: Final = json.loads(request.body) + assert body["model"] == MODEL + assert body["tools"] == list(TOOLS), body + if body["messages"][3]["content"][0]["type"] == "thinking": + assert body["messages"] == list(history), body + assert body["thinking"] == {"type": "enabled", "budget_tokens": 1024}, body + return Reply(status=400, body=SIGNATURE_ERROR) + assert body["messages"] == [*history[:3], tool_use_only_turn, history[4]], body + assert "thinking" not in body, body + return Reply( + body=json.dumps( + { + "id": identity, + "type": "message", + "role": "assistant", + "model": MODEL, + "content": [{"type": "text", "text": "recovered without thinking history"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 30, "output_tokens": 6}, + } + ).encode() + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"anthropic/{MODEL}", api_base=wire.url, api_key=KEY) + response: Final = gateway.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 64, + "thinking": {"type": "enabled", "budget_tokens": 1024}, + "tools": list(TOOLS), + "messages": list(history), + }, + ) + assert response.status_code == 200, response.text + body: Final = response.json() + assert body["id"] == identity, response.text + assert body["content"] == [{"type": "text", "text": "recovered without thinking history"}], response.text + assert body["stop_reason"] == "end_turn", response.text + assert [request.target for request in wire.drain()] == ["/v1/messages", "/v1/messages"] diff --git a/tests/integration/providers/test_anthropic_wire.py b/tests/integration/providers/test_anthropic_wire.py index 27895440712..7e9c5be227a 100644 --- a/tests/integration/providers/test_anthropic_wire.py +++ b/tests/integration/providers/test_anthropic_wire.py @@ -1,4 +1,5 @@ import json +import time import uuid from typing import Final @@ -8,10 +9,17 @@ from integration._support.database import read_rows from integration._support.wire import Reply, Request, wire_server -@pytest.mark.covers("other.provider_wire.anthropic.tool_history_system_cache_and_internal_fields", "quota_management.spend_tracking.cache_tokens.disjoint_classes_use_explicit_rates") +@pytest.mark.covers( + "other.provider_wire.anthropic.tool_history_system_cache_and_internal_fields", + "quota_management.spend_tracking.cache_tokens.disjoint_classes_use_explicit_rates", +) def test_anthropic_tool_history_and_cache_tokens_keep_wire_and_accounting_contracts(gateway: Gateway) -> None: identity: Final = "anthropic-wire-" + uuid.uuid4().hex - tool_schema: Final = {"type": "object", "properties": {"x": {"type": "integer"}, "y": {"type": "integer"}}, "required": ["x", "y"]} + tool_schema: Final = { + "type": "object", + "properties": {"x": {"type": "integer"}, "y": {"type": "integer"}}, + "required": ["x", "y"], + } def respond(request: Request) -> Reply: assert request.method == "POST" and request.target == "/v1/messages" @@ -21,27 +29,80 @@ def test_anthropic_tool_history_and_cache_tokens_keep_wire_and_accounting_contra assert body["system"] == [{"type": "text", "text": "synthetic policy", "cache_control": {"type": "ephemeral"}}] assert body["tools"][0]["name"] == "add" and body["tools"][0]["input_schema"] == tool_schema assert body["max_tokens"] == 16 - assert not {"timeout", "stream_chunk_size", "litellm_params", "litellm_metadata", "rpm", "tpm"}.intersection(body) + assert not {"timeout", "stream_chunk_size", "litellm_params", "litellm_metadata", "rpm", "tpm"}.intersection( + body + ) messages: Final = body["messages"] assert [message["role"] for message in messages] == ["user", "assistant", "user"] assert messages[0]["content"] == [{"type": "text", "text": "first"}] - assert messages[1]["content"] == [{"type": "tool_use", "id": "history-call", "name": "add", "input": {"x": 1, "y": 2}}] - assert messages[2]["content"] == [{"type": "tool_result", "tool_use_id": "history-call", "content": "3"}, {"type": "text", "text": "next"}] - return Reply(body=json.dumps({"id": identity, "type": "message", "role": "assistant", "model": "claude-sonnet-4-5-20250929", "content": [{"type": "tool_use", "id": "next-call", "name": "add", "input": {"x": 3, "y": 4}}], "stop_reason": "tool_use", "stop_sequence": None, "usage": {"input_tokens": 10, "output_tokens": 4, "cache_read_input_tokens": 5, "cache_creation_input_tokens": 7}}).encode()) + assert messages[1]["content"] == [ + {"type": "tool_use", "id": "history-call", "name": "add", "input": {"x": 1, "y": 2}} + ] + assert messages[2]["content"] == [ + {"type": "tool_result", "tool_use_id": "history-call", "content": "3"}, + {"type": "text", "text": "next"}, + ] + return Reply( + body=json.dumps( + { + "id": identity, + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5-20250929", + "content": [{"type": "tool_use", "id": "next-call", "name": "add", "input": {"x": 3, "y": 4}}], + "stop_reason": "tool_use", + "stop_sequence": None, + "usage": { + "input_tokens": 10, + "output_tokens": 4, + "cache_read_input_tokens": 5, + "cache_creation_input_tokens": 7, + }, + } + ).encode() + ) with wire_server(respond) as wire, gateway.scenario() as scenario: - model: Final = scenario.model(model="anthropic/claude-sonnet-4-5-20250929", api_base=wire.url, api_key="synthetic-anthropic-key", input_cost_per_token=0.001, output_cost_per_token=0.002, cache_read_input_token_cost=0.0001, cache_creation_input_token_cost=0.002) - response: Final = gateway.request("POST", "/v1/chat/completions", { - "model": model, "max_tokens": 16, "timeout": 5, - "messages": [ - {"role": "system", "content": [{"type": "text", "text": "synthetic policy", "cache_control": {"type": "ephemeral"}}]}, - {"role": "user", "content": "first"}, - {"role": "assistant", "tool_calls": [{"id": "history-call", "type": "function", "function": {"name": "add", "arguments": '{"x":1,"y":2}'}}]}, - {"role": "tool", "tool_call_id": "history-call", "content": "3"}, - {"role": "user", "content": "next"}, - ], - "tools": [{"type": "function", "function": {"name": "add", "parameters": tool_schema}}], - }) + model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", + api_base=wire.url, + api_key="synthetic-anthropic-key", + input_cost_per_token=0.001, + output_cost_per_token=0.002, + cache_read_input_token_cost=0.0001, + cache_creation_input_token_cost=0.002, + ) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "max_tokens": 16, + "timeout": 5, + "messages": [ + { + "role": "system", + "content": [ + {"type": "text", "text": "synthetic policy", "cache_control": {"type": "ephemeral"}} + ], + }, + {"role": "user", "content": "first"}, + { + "role": "assistant", + "tool_calls": [ + { + "id": "history-call", + "type": "function", + "function": {"name": "add", "arguments": '{"x":1,"y":2}'}, + } + ], + }, + {"role": "tool", "tool_call_id": "history-call", "content": "3"}, + {"role": "user", "content": "next"}, + ], + "tools": [{"type": "function", "function": {"name": "add", "parameters": tool_schema}}], + }, + ) assert response.status_code == 200, response.text body: Final = response.json() assert body["id"].startswith("chatcmpl-") @@ -51,7 +112,14 @@ def test_anthropic_tool_history_and_cache_tokens_keep_wire_and_accounting_contra assert json.loads(tool["function"]["arguments"]) == {"x": 3, "y": 4} assert body["usage"]["prompt_tokens"] == 22 and body["usage"]["completion_tokens"] == 4 assert len(wire.drain()) == 1 - rows: Final = eventually(lambda: read_rows('SELECT spend, prompt_tokens, completion_tokens, metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (body["id"],)), lambda values: len(values) == 1, seconds=70) + rows: Final = eventually( + lambda: read_rows( + 'SELECT spend, prompt_tokens, completion_tokens, metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s', + (body["id"],), + ), + lambda values: len(values) == 1, + seconds=70, + ) assert float(rows[0]["spend"]) == pytest.approx(10 * 0.001 + 5 * 0.0001 + 7 * 0.002 + 4 * 0.002) assert rows[0]["prompt_tokens"] == 22 and rows[0]["completion_tokens"] == 4 metadata: Final = rows[0]["metadata"] @@ -81,3 +149,58 @@ def test_anthropic_bare_string_content_item_is_rejected_as_client_error_before_t ) assert response.status_code == 400, response.text assert wire.drain() == () + + +@pytest.mark.covers("other.provider_wire.anthropic.messages_request_timeout_reaches_transport") +def test_anthropic_messages_slow_upstream_is_cut_off_at_the_deployment_request_timeout(gateway: Gateway) -> None: + identity: Final = "anthropic-timeout-" + uuid.uuid4().hex + prompt: Final = f"slow answer {identity}" + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/messages" + assert request.headers["x-api-key"] == "synthetic-anthropic-key" + body: Final = json.loads(request.body) + assert body["model"] == "claude-sonnet-4-5-20250929" + assert body["max_tokens"] == 16 + assert body["messages"] == [{"role": "user", "content": prompt}] + assert not { + "timeout", + "request_timeout", + "stream_chunk_size", + "litellm_params", + "litellm_metadata", + "rpm", + "tpm", + }.intersection(body) + time.sleep(1.5) + return Reply( + body=json.dumps( + { + "id": identity, + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5-20250929", + "content": [{"type": "text", "text": "late"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 3, "output_tokens": 1}, + } + ).encode() + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", + api_base=wire.url, + api_key="synthetic-anthropic-key", + request_timeout=0.3, + ) + response: Final = gateway.request( + "POST", + "/v1/messages", + {"model": model, "max_tokens": 16, "messages": [{"role": "user", "content": prompt}]}, + headers={"anthropic-version": "2023-06-01"}, + ) + assert response.status_code == 408, response.text + assert "Timeout" in response.json()["error"]["message"], response.text + assert eventually(wire.drain, lambda requests: len(requests) == 1, seconds=5, return_last_on_timeout=True) diff --git a/tests/integration/providers/test_bedrock_converse_client_metadata_wire.py b/tests/integration/providers/test_bedrock_converse_client_metadata_wire.py new file mode 100644 index 00000000000..92d9ebd4a0c --- /dev/null +++ b/tests/integration/providers/test_bedrock_converse_client_metadata_wire.py @@ -0,0 +1,43 @@ +import json +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server +from integration.providers.test_bedrock_auth_wire import MODEL, RESPONSE, TOKEN + +ANTHROPIC_BETA: Final = ["interleaved-thinking-2025-05-14"] +CLIENT_METADATA: Final = {"originator": "codex_cli_rs", "version": "0.1.0", "session_id": "synthetic-session"} + + +def converse_peer(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/model/anthropic.claude-3-haiku-20240307-v1%3A0/converse" + body: Final = json.loads(request.body) + assert body["additionalModelRequestFields"] == {"anthropic_beta": ANTHROPIC_BETA}, body + return Reply(body=RESPONSE) + + +@pytest.mark.covers("providers.bedrock_converse.client_metadata_is_not_forwarded_in_additional_model_request_fields") +def test_client_metadata_is_dropped_from_converse_body_while_anthropic_beta_is_kept(gateway: Gateway) -> None: + with wire_server(converse_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=MODEL, + api_key=TOKEN, + aws_region_name="us-east-1", + aws_bedrock_runtime_endpoint=wire.url, + ) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": "synthetic codex request"}], + "max_tokens": 16, + "anthropic_beta": ANTHROPIC_BETA, + "client_metadata": CLIENT_METADATA, + "cache": {"no-cache": True}, + }, + ) + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"] == "bedrock wire control" + assert len(wire.drain()) == 1, response.text diff --git a/tests/integration/providers/test_bedrock_deepseek_reasoning_wire.py b/tests/integration/providers/test_bedrock_deepseek_reasoning_wire.py new file mode 100644 index 00000000000..f7391da54d9 --- /dev/null +++ b/tests/integration/providers/test_bedrock_deepseek_reasoning_wire.py @@ -0,0 +1,93 @@ +import json +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server + +R1_MODEL: Final = "bedrock/converse/us.deepseek.r1-v1:0" +V3_MODEL: Final = "bedrock/converse/deepseek.v3.2" +TOKEN: Final = "synthetic-bedrock-bearer" +RESPONSE: Final = json.dumps( + { + "output": {"message": {"role": "assistant", "content": [{"text": "deepseek reasoning wire control"}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 9, "outputTokens": 5, "totalTokens": 14}, + "metrics": {"latencyMs": 1}, + } +).encode() + + +def r1_converse_peer(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/model/us.deepseek.r1-v1%3A0/converse", request.target + assert request.headers["authorization"] == f"Bearer {TOKEN}" + body: Final = json.loads(request.body) + assert body["messages"] == [{"role": "user", "content": [{"text": "synthetic r1 request"}]}] + assert body["inferenceConfig"] == {"maxTokens": 16}, body + assert body.get("additionalModelRequestFields") is None, body + return Reply(body=RESPONSE) + + +def v3_converse_peer(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/model/deepseek.v3.2/converse", request.target + assert request.headers["authorization"] == f"Bearer {TOKEN}" + body: Final = json.loads(request.body) + assert body["messages"] == [{"role": "user", "content": [{"text": "synthetic v3 request"}]}] + assert body["inferenceConfig"] == {"maxTokens": 16}, body + assert body["additionalModelRequestFields"] == {"reasoning_effort": "high"}, body + return Reply(body=RESPONSE) + + +@pytest.mark.covers("providers.bedrock_converse.deepseek_r1_drops_thinking_and_reasoning_effort_before_provider") +def test_deepseek_r1_thinking_and_reasoning_effort_are_dropped_instead_of_leaking_into_converse( + gateway: Gateway, +) -> None: + with wire_server(r1_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=R1_MODEL, + api_key=TOKEN, + aws_region_name="us-east-1", + aws_bedrock_runtime_endpoint=wire.url, + drop_params=True, + ) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": "synthetic r1 request"}], + "thinking": {"type": "enabled", "budget_tokens": 1024}, + "reasoning_effort": "high", + "max_tokens": 16, + "cache": {"no-cache": True}, + }, + ) + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"] == "deepseek reasoning wire control", response.text + assert response.json()["usage"]["total_tokens"] == 14, response.text + assert len(wire.drain()) == 1 + + +@pytest.mark.covers("providers.bedrock_converse.deepseek_v3_reasoning_effort_reaches_provider_raw") +def test_deepseek_v3_reasoning_effort_reaches_converse_raw_instead_of_as_anthropic_thinking( + gateway: Gateway, +) -> None: + with wire_server(v3_converse_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=V3_MODEL, api_key=TOKEN, aws_region_name="us-east-1", aws_bedrock_runtime_endpoint=wire.url + ) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": "synthetic v3 request"}], + "reasoning_effort": "high", + "max_tokens": 16, + "cache": {"no-cache": True}, + }, + ) + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"] == "deepseek reasoning wire control", response.text + assert response.json()["usage"]["total_tokens"] == 14, response.text + assert len(wire.drain()) == 1 diff --git a/tests/integration/providers/test_bedrock_invoke_cache_usage_wire.py b/tests/integration/providers/test_bedrock_invoke_cache_usage_wire.py new file mode 100644 index 00000000000..2638c3a2c8d --- /dev/null +++ b/tests/integration/providers/test_bedrock_invoke_cache_usage_wire.py @@ -0,0 +1,86 @@ +import json +from typing import Final + +import pytest +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.wire import Reply, Request, wire_server + +MODEL_ID: Final = "us.amazon.nova-pro-v1:0" +TOKEN: Final = "synthetic-bedrock-bearer" +PROMPT: Final = "summarize the cached policy" +INPUT_TOKENS: Final = 11 +OUTPUT_TOKENS: Final = 4 +CACHE_READ_TOKENS: Final = 900 +CACHE_WRITE_TOKENS: Final = 300 +INPUT_RATE: Final = 0.001 +OUTPUT_RATE: Final = 0.002 +CACHE_READ_RATE: Final = 0.0001 +CACHE_WRITE_RATE: Final = 0.0015 +RESPONSE: Final = json.dumps( + { + "output": {"message": {"role": "assistant", "content": [{"text": "cached policy summary"}]}}, + "stopReason": "end_turn", + "usage": { + "inputTokens": INPUT_TOKENS, + "outputTokens": OUTPUT_TOKENS, + "totalTokens": INPUT_TOKENS + OUTPUT_TOKENS, + "cacheReadInputTokenCount": CACHE_READ_TOKENS, + "cacheWriteInputTokenCount": CACHE_WRITE_TOKENS, + }, + } +).encode() + + +def nova_invoke_peer(request: Request) -> Reply: + assert request.method == "POST" and request.target == f"/model/{MODEL_ID}/invoke", request.target + assert request.headers["authorization"] == f"Bearer {TOKEN}" + body: Final = json.loads(request.body) + assert body["messages"] == [{"role": "user", "content": [{"text": PROMPT}]}], body + return Reply(body=RESPONSE) + + +@pytest.mark.covers("providers.bedrock_invoke.count_suffixed_cache_usage_fields_are_reported_and_charged") +def test_nova_invoke_count_suffixed_cache_usage_fields_are_reported_and_charged(gateway: Gateway) -> None: + with wire_server(nova_invoke_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=f"bedrock/invoke/{MODEL_ID}", + api_key=TOKEN, + aws_region_name="us-east-1", + api_base=wire.url, + input_cost_per_token=INPUT_RATE, + output_cost_per_token=OUTPUT_RATE, + cache_read_input_token_cost=CACHE_READ_RATE, + cache_creation_input_token_cost=CACHE_WRITE_RATE, + ) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "max_tokens": 32, "messages": [{"role": "user", "content": PROMPT}]}, + ) + assert response.status_code == 200, response.text + body: Final = response.json() + assert body["choices"][0]["message"]["content"] == "cached policy summary", response.text + usage: Final = body["usage"] + assert usage["prompt_tokens"] == INPUT_TOKENS + CACHE_READ_TOKENS + CACHE_WRITE_TOKENS, response.text + assert usage["completion_tokens"] == OUTPUT_TOKENS, response.text + assert usage["prompt_tokens_details"]["cached_tokens"] == CACHE_READ_TOKENS, response.text + assert usage["cache_read_input_tokens"] == CACHE_READ_TOKENS, response.text + assert usage["cache_creation_input_tokens"] == CACHE_WRITE_TOKENS, response.text + expected_cost: Final = ( + INPUT_TOKENS * INPUT_RATE + + CACHE_READ_TOKENS * CACHE_READ_RATE + + CACHE_WRITE_TOKENS * CACHE_WRITE_RATE + + OUTPUT_TOKENS * OUTPUT_RATE + ) + assert float(response.headers["x-litellm-response-cost"]) == pytest.approx(expected_cost), response.text + assert len(wire.drain()) == 1 + rows: Final = eventually( + lambda: read_rows( + 'SELECT spend, prompt_tokens FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (body["id"],) + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert float(rows[0]["spend"]) == pytest.approx(expected_cost), rows + assert rows[0]["prompt_tokens"] == INPUT_TOKENS + CACHE_READ_TOKENS + CACHE_WRITE_TOKENS, rows diff --git a/tests/integration/providers/test_bedrock_knowledge_base_user_context_wire.py b/tests/integration/providers/test_bedrock_knowledge_base_user_context_wire.py new file mode 100644 index 00000000000..540b384d694 --- /dev/null +++ b/tests/integration/providers/test_bedrock_knowledge_base_user_context_wire.py @@ -0,0 +1,81 @@ +import json +import uuid +from collections.abc import Callable +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server + +ACCESS_KEY: Final = "AKIAINTEGRATION000003" +USER_CONTEXT: Final = {"userId": "reader@example.com"} +QUERY: Final = "synthetic knowledge base question" +RETRIEVE_RESPONSE: Final = json.dumps( + { + "retrievalResults": [ + { + "content": {"text": "permitted document text"}, + "score": 0.87, + "metadata": { + "x-amz-bedrock-kb-source-uri": "s3://synthetic-bucket/permitted.pdf", + "x-amz-bedrock-kb-chunk-id": "chunk-1", + }, + } + ] + } +).encode() + + +def retrieve_peer(knowledge_base_id: str) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == f"/knowledgebases/{knowledge_base_id}/retrieve", ( + request.target + ) + assert request.headers["authorization"].startswith(f"AWS4-HMAC-SHA256 Credential={ACCESS_KEY}/") + assert json.loads(request.body) == { + "retrievalQuery": {"text": QUERY}, + "retrievalConfiguration": {"vectorSearchConfiguration": {"numberOfResults": 3}}, + "userContext": USER_CONTEXT, + }, request.body + return Reply(body=RETRIEVE_RESPONSE) + + return respond + + +@pytest.mark.covers("providers.bedrock_knowledge_base.search_forwards_user_context_to_retrieve") +def test_vector_store_search_user_context_reaches_bedrock_retrieve_body(gateway: Gateway) -> None: + knowledge_base_id: Final = f"KB{uuid.uuid4().hex[:8].upper()}" + with wire_server(retrieve_peer(knowledge_base_id)) as wire, gateway.scenario() as scenario: + gateway.post( + "/vector_store/new", + { + "vector_store_id": knowledge_base_id, + "custom_llm_provider": "bedrock", + "litellm_params": { + "aws_region_name": "us-east-1", + "aws_access_key_id": ACCESS_KEY, + "aws_secret_access_key": "synthetic-knowledge-base-secret-key", + "aws_bedrock_runtime_endpoint": wire.url, + }, + }, + ) + scenario.cleanups.callback(gateway.post, "/vector_store/delete", {"vector_store_id": knowledge_base_id}) + response: Final = gateway.request( + "POST", + f"/v1/vector_stores/{knowledge_base_id}/search", + {"query": QUERY, "max_num_results": 3, "userContext": USER_CONTEXT}, + ) + assert response.status_code == 200, response.text + assert response.json()["data"] == [ + { + "score": 0.87, + "content": [{"text": "permitted document text", "type": "text"}], + "file_id": "s3://synthetic-bucket/permitted.pdf", + "filename": "permitted.pdf", + "attributes": { + "x-amz-bedrock-kb-source-uri": "s3://synthetic-bucket/permitted.pdf", + "x-amz-bedrock-kb-chunk-id": "chunk-1", + }, + } + ], response.text + assert len(wire.drain()) == 1 diff --git a/tests/integration/providers/test_bedrock_mantle_codex_input_wire.py b/tests/integration/providers/test_bedrock_mantle_codex_input_wire.py index bb7961160dc..aa66e82475b 100644 --- a/tests/integration/providers/test_bedrock_mantle_codex_input_wire.py +++ b/tests/integration/providers/test_bedrock_mantle_codex_input_wire.py @@ -86,3 +86,56 @@ def test_codex_agent_message_context_compaction_and_local_shell_call_reach_mantl forwarded: Final = wire.drain() assert len(forwarded) == 1, forwarded assert JSON_OBJECT.validate_json(forwarded[0].body)["input"] == expected_input, forwarded[0].body + + +SHELL_TOOL: Final[JsonValue] = { + "type": "function", + "name": "shell", + "description": "run a shell command", + "parameters": {"type": "object", "properties": {"command": {"type": "string"}}, "required": ["command"]}, +} +APPLY_PATCH_TOOL: Final[JsonValue] = { + "type": "function", + "name": "apply_patch", + "description": "apply a diff", + "parameters": {"type": "object", "properties": {"patch": {"type": "string"}}, "required": ["patch"]}, +} + + +@pytest.mark.covers("providers.bedrock_mantle.codex_additional_tools_input_item_is_hoisted_to_top_level_tools") +def test_codex_additional_tools_input_item_reaches_mantle_as_top_level_tools(gateway: Gateway) -> None: + marker: Final = uuid.uuid4().hex + expected_input: Final[list[JsonValue]] = [user_turn(f"hoist tools {marker}")] + expected_tools: Final[list[JsonValue]] = [SHELL_TOOL, APPLY_PATCH_TOOL] + + def mantle_peer(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/openai/v1/responses", request.target + assert request.headers["authorization"] == f"Bearer {TOKEN}" + body: Final = JSON_OBJECT.validate_json(request.body) + assert body["model"] == "openai.gpt-5.6-sol", body + assert body["input"] == expected_input, body["input"] + assert body["tools"] == expected_tools, body + return Reply(body=RESPONSE) + + with wire_server(mantle_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=MODEL, api_key=TOKEN, api_base=wire.url, aws_region_name="us-east-2") + response: Final = gateway.request( + "POST", + "/v1/responses", + { + "model": model, + "input": [ + {"type": "additional_tools", "role": "developer", "tools": [APPLY_PATCH_TOOL]}, + user_turn(f"hoist tools {marker}"), + ], + "tools": [SHELL_TOOL], + "store": False, + }, + ) + assert response.status_code == 200, response.text + assert response.json()["output"][0]["content"][0]["text"] == "mantle wire control", response.text + forwarded: Final = wire.drain() + assert len(forwarded) == 1, forwarded + forwarded_body: Final = JSON_OBJECT.validate_json(forwarded[0].body) + assert forwarded_body["input"] == expected_input, forwarded[0].body + assert forwarded_body["tools"] == expected_tools, forwarded[0].body diff --git a/tests/integration/providers/test_bedrock_mantle_responses_wire.py b/tests/integration/providers/test_bedrock_mantle_responses_wire.py index 9bc6f83f8e4..3bd83b5019b 100644 --- a/tests/integration/providers/test_bedrock_mantle_responses_wire.py +++ b/tests/integration/providers/test_bedrock_mantle_responses_wire.py @@ -104,3 +104,43 @@ def test_codex_agent_message_compaction_and_local_shell_items_are_rewritten_for_ } ], response.text assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/openai/v1/responses")] + + +_MANTLE_MIN_MAX_OUTPUT_TOKENS: Final = 16 + + +def _mantle_peer_expecting_max_output_tokens(marker: str, expected: int) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/openai/v1/responses", request.target + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["max_output_tokens"] == expected, request.body.decode() + assert body["input"] == f"clamp probe {marker}", request.body.decode() + return Reply(body=_RESPONSE) + + return respond + + +@pytest.mark.covers("providers.bedrock_mantle.max_output_tokens_below_minimum_is_clamped_to_16_on_the_wire") +def test_max_output_tokens_below_mantle_minimum_is_raised_to_16_before_reaching_mantle(gateway: Gateway) -> None: + marker: Final = uuid4().hex + peer: Final = _mantle_peer_expecting_max_output_tokens(marker, _MANTLE_MIN_MAX_OUTPUT_TOKENS) + with wire_server(peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=_MODEL, api_base=wire.url, api_key=_TOKEN, aws_region_name="us-east-1") + response: Final = gateway.request( + "POST", + "/v1/responses", + {"model": model, "input": f"clamp probe {marker}", "max_output_tokens": 5, "stream": False}, + ) + assert response.status_code == 200, response.text + payload: Final = _JSON_OBJECT.validate_json(response.content) + assert payload["status"] == "completed", response.text + assert payload["output"] == [ + { + **_OUTPUT_MESSAGE, + "phase": None, + "content": [ + {"type": "output_text", "text": "mantle wire control", "annotations": [], "logprobs": None} + ], + } + ], response.text + assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/openai/v1/responses")] diff --git a/tests/integration/providers/test_bedrock_mantle_wire.py b/tests/integration/providers/test_bedrock_mantle_wire.py index 48cd0d770aa..ec32fe5a578 100644 --- a/tests/integration/providers/test_bedrock_mantle_wire.py +++ b/tests/integration/providers/test_bedrock_mantle_wire.py @@ -1,5 +1,7 @@ import json +from collections.abc import Callable from typing import Final +from uuid import uuid4 import pytest from integration._support.client import Gateway @@ -51,3 +53,142 @@ def test_bedrock_mantle_context_overflow_returns_400_saying_prompt_is_too_long(g assert isinstance(message, str), response.text assert f"prompt is too long: {_PROMPT_TOKENS} tokens > {_MODEL_MAXIMUM} maximum" in message, response.text assert [(request.method, request.target) for request in wire.drain()] == [("POST", _RESPONSES_PATH)] + + +_ACCESS_KEY: Final = "AKIAINTEGRATION000003" +_SIGV4_PROMPT: Final = "synthetic sigv4 bridge control" +_SIGV4_RESPONSE: Final = json.dumps( + { + "id": "resp_synthetic_mantle_sigv4", + "object": "response", + "created_at": 1789788253, + "status": "completed", + "model": _BACKEND, + "output": [ + { + "type": "message", + "id": "msg_synthetic_mantle_sigv4", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "mantle sigv4 wire control", "annotations": []}], + } + ], + "usage": {"input_tokens": 21, "output_tokens": 4, "total_tokens": 25}, + } +).encode() + + +def _sigv4_peer(request: Request) -> Reply: + assert request.method == "POST" and request.target == _RESPONSES_PATH, request.target + assert request.headers["authorization"].startswith(f"AWS4-HMAC-SHA256 Credential={_ACCESS_KEY}/"), dict( + request.headers + ) + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["model"] == _BACKEND, body + assert _SIGV4_PROMPT in json.dumps(body["input"]), body + return Reply(body=_SIGV4_RESPONSE) + + +@pytest.mark.covers("providers.bedrock_mantle.chat_bridge_keeps_deployment_aws_credentials_for_sigv4") +def test_chat_completions_bridge_signs_mantle_responses_request_with_deployment_aws_keys(gateway: Gateway) -> None: + with wire_server(_sigv4_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=f"bedrock_mantle/{_BACKEND}", + api_base=wire.url, + api_key=None, + aws_access_key_id=_ACCESS_KEY, + aws_secret_access_key="synthetic-secret-key-for-testing", + aws_region_name="us-east-1", + ) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": _SIGV4_PROMPT}]}, + ) + assert response.status_code == 200, response.text + body: Final = _JSON_OBJECT.validate_json(response.content) + choices: Final = body["choices"] + assert isinstance(choices, list) and len(choices) == 1, response.text + choice: Final = choices[0] + assert isinstance(choice, dict), response.text + assert choice["message"] == {"role": "assistant", "content": "mantle sigv4 wire control"}, response.text + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _RESPONSES_PATH)] + + +_CLAUDE_BACKEND: Final = "anthropic.claude-sonnet-5-v1:0" +_MESSAGES_PATH: Final = "/anthropic/v1/messages" +_STREAM_EVENTS: Final = ( + ( + "message_start", + { + "message": { + "id": "msg_mantle_stream", + "type": "message", + "role": "assistant", + "model": _CLAUDE_BACKEND, + "content": [], + "stop_reason": None, + "stop_sequence": None, + "usage": {"input_tokens": 11, "output_tokens": 1}, + } + }, + ), + ("content_block_start", {"index": 0, "content_block": {"type": "text", "text": ""}}), + ("content_block_delta", {"index": 0, "delta": {"type": "text_delta", "text": "mantle "}}), + ("content_block_delta", {"index": 0, "delta": {"type": "text_delta", "text": "stream control"}}), + ("content_block_stop", {"index": 0}), + ("message_delta", {"delta": {"stop_reason": "end_turn", "stop_sequence": None}, "usage": {"output_tokens": 4}}), + ("message_stop", {}), +) +_STREAM_FRAMES: Final = tuple( + f"event: {kind}\ndata: {json.dumps({'type': kind, **payload})}\n\n".encode() for kind, payload in _STREAM_EVENTS +) + + +def _streaming_messages_peer(prompt: str) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + assert request.method == "POST" + assert request.target == _MESSAGES_PATH + assert request.headers["authorization"] == f"Bearer {_API_KEY}" + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["model"] == _CLAUDE_BACKEND, body + assert body["stream"] is True, body + assert body["messages"] == [{"role": "user", "content": prompt}], body + return Reply(content_type="text/event-stream", chunks=_STREAM_FRAMES) + + return respond + + +@pytest.mark.covers("providers.bedrock_mantle.messages_stream_sends_stream_true_and_relays_sse_events") +def test_bedrock_mantle_messages_stream_relays_anthropic_sse_instead_of_failing_on_event_stream_decode( + gateway: Gateway, +) -> None: + prompt: Final = f"synthetic mantle stream control {uuid4().hex}" + with wire_server(_streaming_messages_peer(prompt)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=f"bedrock_mantle/{_CLAUDE_BACKEND}", api_base=wire.url, api_key=_API_KEY, aws_region_name="us-east-1" + ) + with gateway.client.stream( + "POST", + "/v1/messages", + json={ + "model": model, + "max_tokens": 64, + "stream": True, + "messages": [{"role": "user", "content": prompt}], + }, + headers={"Authorization": f"Bearer {gateway.key}"}, + ) as response: + assert response.status_code == 200, response.read().decode() + assert response.headers["content-type"].startswith("text/event-stream"), dict(response.headers) + events: Final = tuple( + _JSON_OBJECT.validate_json(line.removeprefix("data: ")) + for line in response.iter_lines() + if line.startswith("data: ") + ) + assert tuple(event["type"] for event in events) == tuple(kind for kind, _ in _STREAM_EVENTS), events + assert ( + "".join(str(event["delta"]["text"]) for event in events if event["type"] == "content_block_delta") + == "mantle stream control" + ), events + assert [(request.method, request.target) for request in wire.drain()] == [("POST", _MESSAGES_PATH)] diff --git a/tests/integration/providers/test_bedrock_marengo_embed_3_wire.py b/tests/integration/providers/test_bedrock_marengo_embed_3_wire.py new file mode 100644 index 00000000000..9cdedef6ef6 --- /dev/null +++ b/tests/integration/providers/test_bedrock_marengo_embed_3_wire.py @@ -0,0 +1,35 @@ +import json +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server + +MODEL: Final = "bedrock/us.twelvelabs.marengo-embed-3-0-v1:0" +TOKEN: Final = "synthetic-bedrock-bearer" +INPUT: Final = "hello world" +VECTOR: Final = [0.1, 0.2, 0.3] +RESPONSE: Final = json.dumps({"data": [{"embedding": VECTOR}]}).encode() + + +def marengo_3_peer(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.target == "/model/us.twelvelabs.marengo-embed-3-0-v1%3A0/invoke", request.target + assert request.headers["authorization"] == f"Bearer {TOKEN}" + assert json.loads(request.body) == {"inputType": "text", "text": {"inputText": INPUT}}, request.body + return Reply(body=RESPONSE) + + +@pytest.mark.covers("providers.bedrock_embedding.marengo_3_text_input_reaches_bedrock_nested_under_input_type") +def test_marengo_3_text_embedding_nests_input_text_under_input_type(gateway: Gateway) -> None: + with wire_server(marengo_3_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=MODEL, + api_key=TOKEN, + api_base=wire.url, + aws_region_name="us-east-1", + ) + response: Final = gateway.request("POST", "/v1/embeddings", {"model": model, "input": INPUT}) + assert response.status_code == 200, response.text + assert response.json()["data"] == [{"object": "embedding", "index": 0, "embedding": VECTOR}], response.text + assert len(wire.drain()) == 1, "the embedding request never reached Bedrock" diff --git a/tests/integration/providers/test_bedrock_role_configuration.py b/tests/integration/providers/test_bedrock_role_configuration.py index ac8edbdfde0..857535e5e35 100644 --- a/tests/integration/providers/test_bedrock_role_configuration.py +++ b/tests/integration/providers/test_bedrock_role_configuration.py @@ -7,7 +7,6 @@ from urllib.parse import parse_qs import pytest import yaml - from integration._support.client import Gateway from integration._support.process import owned_proxy from integration._support.wire import Reply, Request, wire_server @@ -30,7 +29,10 @@ def test_role_reference_from_db_and_yaml_reaches_real_sts_http_and_bedrock(gatew assert parameters["RoleArn"] == [role] assert parameters["RoleSessionName"][0] in {"integration-yaml-session", "integration-db-session"} result = f"{assumed_key}synthetic-assumed-secret-key-for-testing{assumed_token}2035-01-01T00:00:00Zarn:aws:sts::123456789012:assumed-role/integration/sessionintegration:session0" - return Reply(content_type="text/xml", body=f'<{action}Response xmlns="https://sts.amazonaws.com/doc/2011-06-15/">{result}synthetic-sts-request'.encode()) + return Reply( + content_type="text/xml", + body=f'<{action}Response xmlns="https://sts.amazonaws.com/doc/2011-06-15/">{result}synthetic-sts-request'.encode(), + ) def bedrock(request: Request) -> Reply: assert request.method == "POST" and request.target == "/model/anthropic.claude-3-haiku-20240307-v1%3A0/converse" @@ -41,8 +43,11 @@ def test_role_reference_from_db_and_yaml_reaches_real_sts_http_and_bedrock(gatew with wire_server(sts) as authority, wire_server(bedrock) as provider: parameters: Final = { - "model": MODEL, "aws_region_name": "us-east-1", "aws_role_name": "os.environ/INTEGRATION_ROLE_ARN", - "aws_session_name": "integration-yaml-session", "aws_bedrock_runtime_endpoint": provider.url, + "model": MODEL, + "aws_region_name": "us-east-1", + "aws_role_name": "os.environ/INTEGRATION_ROLE_ARN", + "aws_session_name": "integration-yaml-session", + "aws_bedrock_runtime_endpoint": provider.url, "aws_sts_endpoint": authority.url, } alias: Final = "integration-role-yaml-" + uuid.uuid4().hex @@ -53,23 +58,144 @@ def test_role_reference_from_db_and_yaml_reaches_real_sts_http_and_bedrock(gatew empty: Final = tmp_path / "empty-aws-config" empty.write_text("") overrides: Final = { - "INTEGRATION_ROLE_ARN": role, "AWS_ACCESS_KEY_ID": "AKIAINTEGRATION000001", "AWS_SECRET_ACCESS_KEY": "synthetic-source-secret-key-for-testing", - "AWS_CONFIG_FILE": str(empty), "AWS_SHARED_CREDENTIALS_FILE": str(empty), "AWS_EC2_METADATA_DISABLED": "true", - "AWS_ENDPOINT_URL_STS": authority.url, "AWS_DEFAULT_REGION": "us-east-1", "LITELLM_RUST": "false", + "INTEGRATION_ROLE_ARN": role, + "AWS_ACCESS_KEY_ID": "AKIAINTEGRATION000001", + "AWS_SECRET_ACCESS_KEY": "synthetic-source-secret-key-for-testing", + "AWS_CONFIG_FILE": str(empty), + "AWS_SHARED_CREDENTIALS_FILE": str(empty), + "AWS_EC2_METADATA_DISABLED": "true", + "AWS_ENDPOINT_URL_STS": authority.url, + "AWS_DEFAULT_REGION": "us-east-1", + "LITELLM_RUST": "false", } - with owned_proxy(gateway, tmp_path, overrides, config=path, remove_environment=tuple(name for name in os.environ if name.startswith("AWS_"))) as candidate, candidate.scenario() as scenario: - database_model: Final = scenario.model(**{**parameters, "api_key": None, "aws_session_name": "integration-db-session"}) + with ( + owned_proxy( + gateway, + tmp_path, + overrides, + config=path, + remove_environment=tuple(name for name in os.environ if name.startswith("AWS_")), + ) as candidate, + candidate.scenario() as scenario, + ): + database_model: Final = scenario.model( + **{**parameters, "api_key": None, "aws_session_name": "integration-db-session"} + ) for generation in range(2): for model in (alias, database_model): - response: Final = candidate.request("POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": "synthetic role request"}], "cache": {"no-cache": True}}) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": "synthetic role request"}], + "cache": {"no-cache": True}, + }, + ) assert response.status_code == 200, response.text assert response.json()["choices"][0]["message"]["content"] == "bedrock wire control" assert response.json()["usage"]["total_tokens"] == 15 assert len(provider.drain()) == 1 if generation == 0: - target: Final = next(entry for entry in candidate.get("/model/info")["data"] if entry["model_name"] == database_model) - response: Final = candidate.request("PATCH", f"/model/{target['model_info']['id']}/update", {"model_info": {"description": "role reload"}}) + target: Final = next( + entry for entry in candidate.get("/model/info")["data"] if entry["model_name"] == database_model + ) + response: Final = candidate.request( + "PATCH", + f"/model/{target['model_info']['id']}/update", + {"model_info": {"description": "role reload"}}, + ) assert response.status_code == 200, response.text - assumed: Final = tuple(parse_qs(request.body.decode()) for request in authority.drain() if parse_qs(request.body.decode())["Action"] == ["AssumeRole"]) - assert {entry["RoleSessionName"][0] for entry in assumed} == {"integration-yaml-session", "integration-db-session"} + assumed: Final = tuple( + parse_qs(request.body.decode()) + for request in authority.drain() + if parse_qs(request.body.decode())["Action"] == ["AssumeRole"] + ) + assert {entry["RoleSessionName"][0] for entry in assumed} == { + "integration-yaml-session", + "integration-db-session", + } assert all(entry["RoleArn"] == [role] for entry in assumed) + + +@pytest.mark.covers("providers.bedrock_assume_role.repeat_requests_reuse_cached_sts_session_per_session_name") +def test_repeat_requests_under_one_session_name_assume_role_once_per_session_name( + gateway: Gateway, tmp_path: Path +) -> None: + role: Final = "arn:aws:iam::123456789012:role/integration-" + uuid.uuid4().hex + assumed_key: Final = "ASIAINTEGRATION000002" + assumed_token: Final = "synthetic-cached-session-token" + first_session: Final = "integration-attributed-user-a-" + uuid.uuid4().hex[:8] + second_session: Final = "integration-attributed-user-b-" + uuid.uuid4().hex[:8] + + def sts(request: Request) -> Reply: + parameters: Final = parse_qs(request.body.decode()) + action: Final = parameters["Action"][0] + assert request.method == "POST" and action in {"GetCallerIdentity", "AssumeRole"} + if action == "GetCallerIdentity": + result = "arn:aws:iam::123456789012:user/integration-sourceintegration-source123456789012" + else: + assert parameters["RoleArn"] == [role] + result = f"{assumed_key}synthetic-assumed-secret-key-for-testing{assumed_token}2035-01-01T00:00:00Zarn:aws:sts::123456789012:assumed-role/integration/sessionintegration:session0" + return Reply( + content_type="text/xml", + body=f'<{action}Response xmlns="https://sts.amazonaws.com/doc/2011-06-15/">{result}synthetic-sts-request'.encode(), + ) + + def bedrock(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/model/anthropic.claude-3-haiku-20240307-v1%3A0/converse" + assert f"Credential={assumed_key}/" in request.headers["authorization"] + assert request.headers["x-amz-security-token"] == assumed_token + return Reply(body=RESPONSE) + + with wire_server(sts) as authority, wire_server(bedrock) as provider: + empty: Final = tmp_path / "empty-aws-config" + empty.write_text("") + overrides: Final = { + "AWS_ACCESS_KEY_ID": "AKIAINTEGRATION000002", + "AWS_SECRET_ACCESS_KEY": "synthetic-source-secret-key-for-testing", + "AWS_CONFIG_FILE": str(empty), + "AWS_SHARED_CREDENTIALS_FILE": str(empty), + "AWS_EC2_METADATA_DISABLED": "true", + "AWS_ENDPOINT_URL_STS": authority.url, + "AWS_DEFAULT_REGION": "us-east-1", + "LITELLM_RUST": "false", + } + with ( + owned_proxy( + gateway, + tmp_path, + overrides, + remove_environment=tuple(name for name in os.environ if name.startswith("AWS_")), + ) as candidate, + candidate.scenario() as scenario, + ): + parameters: Final = { + "model": MODEL, + "api_key": None, + "aws_region_name": "us-east-1", + "aws_role_name": role, + "aws_bedrock_runtime_endpoint": provider.url, + "aws_sts_endpoint": authority.url, + } + first_model: Final = scenario.model(**{**parameters, "aws_session_name": first_session}) + second_model: Final = scenario.model(**{**parameters, "aws_session_name": second_session}) + for model in (first_model, first_model, second_model, second_model): + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": "synthetic cached role request"}], + "cache": {"no-cache": True}, + }, + ) + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"] == "bedrock wire control" + assert len(provider.drain()) == 1 + assumed: Final = tuple( + parse_qs(request.body.decode()) + for request in authority.drain() + if parse_qs(request.body.decode())["Action"] == ["AssumeRole"] + ) + assert tuple(entry["RoleSessionName"][0] for entry in assumed) == (first_session, second_session), assumed diff --git a/tests/integration/providers/test_bedrock_thinking_tokens_wire.py b/tests/integration/providers/test_bedrock_thinking_tokens_wire.py index 074adeb41f6..19d7b1e291c 100644 --- a/tests/integration/providers/test_bedrock_thinking_tokens_wire.py +++ b/tests/integration/providers/test_bedrock_thinking_tokens_wire.py @@ -1,4 +1,5 @@ import json +import uuid from typing import Final import pytest @@ -34,13 +35,13 @@ _JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) _JSON_LIST: Final = TypeAdapter(list[dict[str, JsonValue]]) -def redacted_thinking_peer(request: Request) -> Reply: +def redacted_thinking_peer(request: Request, prompts: tuple[str, str]) -> Reply: assert request.method == "POST" and request.target == "/model/global.anthropic.claude-opus-4-8/converse" assert request.headers["authorization"] == f"Bearer {TOKEN}" body: Final = json.loads(request.body) assert body["messages"] in ( - [{"role": "user", "content": [{"text": PROMPT}]}], - [{"role": "user", "content": [{"text": RESPONSES_PROMPT}]}], + [{"role": "user", "content": [{"text": prompts[0]}]}], + [{"role": "user", "content": [{"text": prompts[1]}]}], ), body assert body["additionalModelRequestFields"]["thinking"]["type"] == "adaptive", body return Reply(body=RESPONSE) @@ -48,7 +49,9 @@ def redacted_thinking_peer(request: Request) -> Reply: @pytest.mark.covers("other.provider_wire.bedrock.hidden_thinking_tokens_are_not_reported_as_text") def test_bedrock_redacted_thinking_is_not_reported_as_zero_reasoning_tokens(gateway: Gateway) -> None: - with wire_server(redacted_thinking_peer) as wire, gateway.scenario() as scenario: + identity: Final = " " + uuid.uuid4().hex + prompts: Final = (PROMPT + identity, RESPONSES_PROMPT + identity) + with wire_server(lambda request: redacted_thinking_peer(request, prompts)) as wire, gateway.scenario() as scenario: model: Final = scenario.model( model=MODEL, api_key=TOKEN, aws_region_name="us-east-1", aws_bedrock_runtime_endpoint=wire.url ) @@ -57,7 +60,7 @@ def test_bedrock_redacted_thinking_is_not_reported_as_zero_reasoning_tokens(gate "/v1/chat/completions", { "model": model, - "messages": [{"role": "user", "content": PROMPT}], + "messages": [{"role": "user", "content": prompts[0]}], "max_tokens": 4000, "reasoning_effort": "max", }, @@ -76,7 +79,7 @@ def test_bedrock_redacted_thinking_is_not_reported_as_zero_reasoning_tokens(gate responses: Final = gateway.request( "POST", "/v1/responses", - {"model": model, "input": RESPONSES_PROMPT, "max_output_tokens": 4000, "reasoning": {"effort": "max"}}, + {"model": model, "input": prompts[1], "max_output_tokens": 4000, "reasoning": {"effort": "max"}}, ) assert responses.status_code == 200, responses.text responses_body: Final = _JSON_OBJECT.validate_json(responses.content) diff --git a/tests/integration/providers/test_fireworks_ai_session_affinity_wire.py b/tests/integration/providers/test_fireworks_ai_session_affinity_wire.py new file mode 100644 index 00000000000..9c8490a4591 --- /dev/null +++ b/tests/integration/providers/test_fireworks_ai_session_affinity_wire.py @@ -0,0 +1,69 @@ +import json +import uuid +from typing import Final + +import pytest +from integration._support.client import Gateway, eventually, object_value +from integration._support.database import read_rows +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_MODEL: Final = "accounts/fireworks/models/kimi-k3" +_API_KEY: Final = "synthetic-fireworks-key" +_PROMPT: Final = "keep this conversation on one replica" +_SESSION_ID: Final = "conversation-affinity-6220" +_CACHED_TOKENS: Final = 7 +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + + +def _cached_reply(request: Request, identity: str) -> Reply: + assert request.method == "POST" + assert request.target == "/chat/completions" + assert request.headers["authorization"] == f"Bearer {_API_KEY}" + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["model"] == _MODEL, body + return Reply( + body=json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": _MODEL, + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "pinned"}, "finish_reason": "stop"} + ], + "usage": { + "prompt_tokens": 12, + "completion_tokens": 1, + "total_tokens": 13, + "prompt_tokens_details": {"cached_tokens": _CACHED_TOKENS}, + }, + } + ).encode() + ) + + +@pytest.mark.covers("other.provider_wire.fireworks_ai.session_id_sent_as_affinity_header_and_cached_tokens_logged") +def test_fireworks_session_id_sends_affinity_header_and_logs_cache_read_tokens(gateway: Gateway) -> None: + identity: Final = f"fw-session-affinity-{uuid.uuid4().hex}" + with wire_server(lambda request: _cached_reply(request, identity)) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"fireworks_ai/{_MODEL}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": _PROMPT}]}, + headers={"x-litellm-session-id": _SESSION_ID}, + ) + assert response.status_code == 200, response.text + payload: Final = _JSON_OBJECT.validate_json(response.content) + assert payload["id"] == identity, response.text + requests: Final = wire.drain() + assert [(request.method, request.target) for request in requests] == [("POST", "/chat/completions")] + assert requests[0].headers.get("x-session-affinity") == _SESSION_ID, requests[0].headers + rows: Final = eventually( + lambda: read_rows('SELECT metadata FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (identity,)), + lambda values: len(values) == 1, + seconds=70, + ) + usage_values: Final = object_value(object_value(rows[0]["metadata"])["additional_usage_values"]) + assert usage_values.get("cache_read_input_tokens") == _CACHED_TOKENS, rows[0]["metadata"] diff --git a/tests/integration/providers/test_gemini_messages_cache_control_wire.py b/tests/integration/providers/test_gemini_messages_cache_control_wire.py new file mode 100644 index 00000000000..71073b3dd5e --- /dev/null +++ b/tests/integration/providers/test_gemini_messages_cache_control_wire.py @@ -0,0 +1,86 @@ +import json +import uuid +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "gemini-2.5-flash" +_API_KEY: Final = "synthetic-gemini-key" +_CACHE_NAME: Final = "cachedContents/synthetic-cache" +_CACHED_POLICY: Final = " ".join(f"policy clause {index} applies" for index in range(600)) +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + + +def _generate_content_reply(text: str) -> bytes: + return json.dumps( + { + "candidates": [ + {"content": {"parts": [{"text": text}], "role": "model"}, "finishReason": "STOP", "index": 0} + ], + "usageMetadata": { + "promptTokenCount": 1300, + "candidatesTokenCount": 5, + "totalTokenCount": 1305, + "cachedContentTokenCount": 1290, + }, + "modelVersion": _BACKEND, + } + ).encode() + + +@pytest.mark.covers("other.provider_wire.gemini.messages_cache_control_creates_cached_content_with_anthropic_ttl") +def test_gemini_messages_cache_control_creates_cached_content_and_generates_from_it(gateway: Gateway) -> None: + identity: Final = f"gemini-messages-cache-{uuid.uuid4().hex}" + user_prompt: Final = f"Summarize the policy. Request {identity}." + + def respond(request: Request) -> Reply: + assert request.headers["x-goog-api-key"] == _API_KEY, request.headers + if request.method == "GET": + assert request.target == f"/models/{_BACKEND}:cachedContents", request.target + return Reply(body=b"{}") + assert request.method == "POST", request.method + body: Final = _JSON_OBJECT.validate_json(request.body) + if request.target == f"/models/{_BACKEND}:cachedContents": + assert isinstance(body["displayName"], str) and body["displayName"], body + assert body == { + "contents": [{"role": "user", "parts": [{"text": "."}]}], + "model": f"models/{_BACKEND}", + "displayName": body["displayName"], + "ttl": "300s", + "system_instruction": {"parts": [{"text": _CACHED_POLICY}]}, + "tools": None, + } + return Reply(body=json.dumps({"name": _CACHE_NAME, "model": f"models/{_BACKEND}"}).encode()) + assert request.target == f"/models/{_BACKEND}:generateContent", request.target + assert body == { + "contents": [{"role": "user", "parts": [{"text": user_prompt}]}], + "generationConfig": {"max_output_tokens": 32}, + "cachedContent": _CACHE_NAME, + } + return Reply(body=_generate_content_reply("The policy applies.")) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=f"gemini/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + response: Final = gateway.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 32, + "system": [ + {"type": "text", "text": _CACHED_POLICY, "cache_control": {"type": "ephemeral", "ttl": "5m"}} + ], + "messages": [{"role": "user", "content": user_prompt}], + }, + ) + assert response.status_code == 200, response.text + payload: Final = _JSON_OBJECT.validate_json(response.content) + assert payload["content"] == [{"type": "text", "text": "The policy applies."}], response.text + assert [(request.method, request.target) for request in wire.drain()] == [ + ("GET", f"/models/{_BACKEND}:cachedContents"), + ("POST", f"/models/{_BACKEND}:cachedContents"), + ("POST", f"/models/{_BACKEND}:generateContent"), + ] diff --git a/tests/integration/providers/test_nvidia_nim_ranking_wire.py b/tests/integration/providers/test_nvidia_nim_ranking_wire.py new file mode 100644 index 00000000000..9ed7a4eb48e --- /dev/null +++ b/tests/integration/providers/test_nvidia_nim_ranking_wire.py @@ -0,0 +1,47 @@ +import json +from typing import Final + +import pytest +from integration._support.client import JSON_OBJECT, Gateway +from integration._support.wire import Reply, Request, wire_server + +MODEL: Final = "nvidia_nim/ranking/nvidia/llama-3.2-nv-rerankqa-1b-v2" +QUERY: Final = "which passage shows the gateway diagram" +IMAGE_PASSAGE: Final = "data:image/png;base64,aW50ZWdyYXRpb24tc3ludGhldGljLWltYWdl" +TEXT_PASSAGE: Final = "the gateway proxies rerank calls" +RESPONSE: Final = json.dumps( + {"rankings": [{"index": 0, "logit": 0.82}, {"index": 1, "logit": -1.4}], "usage": {"total_tokens": 11}} +).encode() + + +def ranking_peer(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/ranking", request.target + assert request.headers["authorization"] == "Bearer integration-provider-key" + body: Final = JSON_OBJECT.validate_json(request.body) + assert body == { + "model": "nvidia/llama-3.2-nv-rerankqa-1b-v2", + "query": {"text": QUERY}, + "passages": [{"image": IMAGE_PASSAGE}, {"text": TEXT_PASSAGE}], + }, body + return Reply(body=RESPONSE) + + +@pytest.mark.covers( + "providers.nvidia_nim_ranking.image_passages_reach_ranking_without_top_k_and_top_n_is_applied_locally" +) +def test_nvidia_nim_ranking_keeps_image_passages_and_applies_top_n_without_sending_top_k(gateway: Gateway) -> None: + with wire_server(ranking_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model=MODEL, api_base=wire.url, model_info={"mode": "rerank"}) + response: Final = gateway.request( + "POST", + "/v1/rerank", + { + "model": model, + "query": QUERY, + "documents": [{"image": IMAGE_PASSAGE}, {"text": TEXT_PASSAGE}], + "top_n": 1, + }, + ) + assert response.status_code == 200, response.text + assert response.json()["results"] == [{"index": 0, "relevance_score": 0.82}], response.text + assert len(wire.drain()) == 1, "Expected exactly one provider ranking call" diff --git a/tests/integration/providers/test_rerank_latency_headers_wire.py b/tests/integration/providers/test_rerank_latency_headers_wire.py new file mode 100644 index 00000000000..62624e0480a --- /dev/null +++ b/tests/integration/providers/test_rerank_latency_headers_wire.py @@ -0,0 +1,50 @@ +import json +import uuid +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server + +MODEL: Final = "cohere/synthetic-rerank-model-without-pricing" +QUERY: Final = "which document mentions the gateway" +DOCUMENTS: Final = ("the gateway proxies rerank calls", "unrelated synthetic text") +RESPONSE: Final = json.dumps( + { + "id": "synthetic-rerank-id", + "results": [{"index": 0, "relevance_score": 0.91}, {"index": 1, "relevance_score": 0.03}], + "meta": {"api_version": {"version": "2"}, "billed_units": {"search_units": 1}}, + } +).encode() + + +def rerank_peer(request: Request) -> Reply: + assert request.method == "POST" and request.target.endswith("/rerank"), request.target + body: Final = json.loads(request.body) + assert body["query"] == QUERY and body["documents"] == list(DOCUMENTS), request.body + return Reply(body=RESPONSE) + + +@pytest.mark.covers("providers.rerank.response_carries_latency_and_cost_headers") +def test_rerank_response_carries_call_id_latency_and_cost_headers_like_chat_completions(gateway: Gateway) -> None: + with wire_server(rerank_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=MODEL, + api_key="synthetic-cohere-key", + api_base=wire.url, + model_info={"mode": "rerank"}, + ) + response: Final = gateway.request( + "POST", "/v1/rerank", {"model": model, "query": QUERY, "documents": list(DOCUMENTS), "top_n": 2} + ) + assert response.status_code == 200, response.text + assert [(result["index"], result["relevance_score"]) for result in response.json()["results"]] == [ + (0, 0.91), + (1, 0.03), + ], response.text + assert len(wire.drain()) == 1, "Expected exactly one provider rerank call" + assert response.headers["x-litellm-model-group"] == model, response.text + assert uuid.UUID(response.headers["x-litellm-call-id"]).version == 4, response.headers + assert float(response.headers["x-litellm-response-cost"]) == 0.0, response.headers + assert float(response.headers["x-litellm-response-duration-ms"]) > 0, response.headers + assert float(response.headers["x-litellm-overhead-duration-ms"]) >= 0, response.headers diff --git a/tests/integration/providers/test_responses_bridge_incomplete.py b/tests/integration/providers/test_responses_bridge_incomplete.py index 3352ca5775e..e700d17ea88 100644 --- a/tests/integration/providers/test_responses_bridge_incomplete.py +++ b/tests/integration/providers/test_responses_bridge_incomplete.py @@ -62,3 +62,129 @@ def test_chat_over_responses_deployment_returns_length_when_output_tokens_run_ou assert body["choices"][0]["message"]["role"] == "assistant", response.text assert body["usage"]["prompt_tokens"] == 12 and body["usage"]["completion_tokens"] == 16, response.text assert body["usage"]["total_tokens"] == 28, response.text + + +@pytest.mark.covers("other.provider_wire.responses_bridge.sub_minimum_max_tokens_clamped_to_provider_floor") +def test_messages_over_responses_deployment_with_max_tokens_1_is_clamped_to_16_instead_of_400(gateway: Gateway) -> None: + identity: Final = "responses-clamp-" + uuid.uuid4().hex + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/responses", request.target + assert request.headers["authorization"] == "Bearer synthetic-openai-key" + body: Final = json.loads(request.body) + assert body["model"] == "gpt-5.4" + if body["max_output_tokens"] < 16: + return Reply( + status=400, + body=json.dumps( + { + "error": { + "message": "Invalid 'max_output_tokens': integer below minimum value. Expected a value >= 16, but got 1 instead.", + "type": "invalid_request_error", + "param": "max_output_tokens", + "code": "integer_below_min_value", + } + } + ).encode(), + ) + assert body["max_output_tokens"] == 16 + assert body["input"] == [ + { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": f"warmup probe {identity}"}], + } + ] + return Reply( + body=json.dumps( + { + "id": f"resp_{identity}", + "object": "response", + "created_at": 1789788253, + "status": "completed", + "model": "gpt-5.4", + "output": [ + { + "type": "message", + "id": f"msg_{identity}", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "ok", "annotations": []}], + } + ], + "usage": {"input_tokens": 12, "output_tokens": 1, "total_tokens": 13}, + } + ).encode() + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/responses/gpt-5.4", api_base=wire.url, api_key="synthetic-openai-key" + ) + response: Final = gateway.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 1, + "messages": [{"role": "user", "content": f"warmup probe {identity}"}], + }, + ) + assert response.status_code == 200, response.text + body: Final = response.json() + assert len(wire.drain()) == 1 + assert body["role"] == "assistant", response.text + assert body["content"] == [{"type": "text", "text": "ok"}], response.text + assert body["stop_reason"] == "end_turn", response.text + + +@pytest.mark.covers("providers.responses_bridge.sub_minimum_max_tokens_is_raised_to_the_openai_floor") +def test_messages_over_responses_deployment_with_max_tokens_one_reaches_openai_as_sixteen(gateway: Gateway) -> None: + identity: Final = "responses-min-tokens-" + uuid.uuid4().hex + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/responses", request.target + assert request.headers["authorization"] == "Bearer synthetic-openai-key" + body: Final = json.loads(request.body) + assert body["model"] == "gpt-5.6-sol" + assert body["max_output_tokens"] == 16, body + return Reply( + body=json.dumps( + { + "id": f"resp_{identity}", + "object": "response", + "created_at": 1789788253, + "status": "completed", + "model": "gpt-5.6-sol", + "output": [ + { + "type": "message", + "id": f"msg_{identity}", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "ok", "annotations": []}], + } + ], + "usage": {"input_tokens": 9, "output_tokens": 1, "total_tokens": 10}, + } + ).encode() + ) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model="openai/responses/gpt-5.6-sol", api_base=wire.url, api_key="synthetic-openai-key" + ) + response: Final = gateway.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 1, + "messages": [{"role": "user", "content": f"warmup {identity}"}], + }, + ) + assert response.status_code == 200, response.text + body: Final = response.json() + assert len(wire.drain()) == 1 + assert body["content"] == [{"type": "text", "text": "ok"}], response.text + assert body["usage"]["input_tokens"] == 9 and body["usage"]["output_tokens"] == 1, response.text diff --git a/tests/integration/providers/test_responses_bridge_namespace_tools.py b/tests/integration/providers/test_responses_bridge_namespace_tools.py new file mode 100644 index 00000000000..746dac03ced --- /dev/null +++ b/tests/integration/providers/test_responses_bridge_namespace_tools.py @@ -0,0 +1,159 @@ +import json +import uuid +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +JSON_LIST: Final = TypeAdapter(list[dict[str, JsonValue]]) +NAMESPACE: Final = "mcp__everything" +TOOL_NAME: Final = "get_sum" +FLATTENED_NAME: Final = f"{NAMESPACE}__{TOOL_NAME}" +CALL_ID: Final = "call_synthetic_get_sum" +ARGUMENTS: Final = json.dumps({"a": 2, "b": 3}) +PARAMETERS: Final[dict[str, JsonValue]] = { + "type": "object", + "required": ["a", "b"], + "properties": {"a": {"type": "number"}, "b": {"type": "number"}}, +} +NAMESPACE_TOOL: Final[dict[str, JsonValue]] = { + "type": "namespace", + "name": NAMESPACE, + "description": "Tools exposed by the everything MCP server", + "tools": [ + { + "type": "function", + "name": TOOL_NAME, + "description": "Adds two numbers", + "strict": False, + "parameters": PARAMETERS, + } + ], +} +EXPECTED_CHAT_TOOLS: Final[list[JsonValue]] = [ + { + "type": "function", + "function": { + "name": FLATTENED_NAME, + "description": "Tools exposed by the everything MCP server\n\nAdds two numbers", + "parameters": PARAMETERS, + "strict": False, + }, + } +] + + +def tool_call_completion(marker: str) -> bytes: + return json.dumps( + { + "id": f"chatcmpl-{marker}", + "object": "chat.completion", + "created": 1789788253, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "finish_reason": "tool_calls", + "message": { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": CALL_ID, + "type": "function", + "function": {"name": FLATTENED_NAME, "arguments": ARGUMENTS}, + } + ], + }, + } + ], + "usage": {"prompt_tokens": 30, "completion_tokens": 12, "total_tokens": 42}, + } + ).encode() + + +def text_completion(marker: str) -> bytes: + return json.dumps( + { + "id": f"chatcmpl-{marker}-final", + "object": "chat.completion", + "created": 1789788254, + "model": "gpt-4o-mini", + "choices": [ + { + "index": 0, + "finish_reason": "stop", + "message": {"role": "assistant", "content": "The sum is 5"}, + } + ], + "usage": {"prompt_tokens": 40, "completion_tokens": 5, "total_tokens": 45}, + } + ).encode() + + +@pytest.mark.covers("other.provider_wire.responses_bridge.codex_namespace_tools_reach_chat_upstream_and_round_trip") +def test_codex_namespace_tool_is_flattened_for_chat_upstream_and_restored_in_responses_output( + gateway: Gateway, +) -> None: + marker: Final = uuid.uuid4().hex + prompt: Final = f"add 2 and 3 {marker}" + + def chat_peer(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/chat/completions", request.target + body: Final = JSON_OBJECT.validate_json(request.body) + assert body["tools"] == EXPECTED_CHAT_TOOLS, body + messages: Final = JSON_LIST.validate_python(body["messages"]) + if len(messages) == 1: + return Reply(body=tool_call_completion(marker)) + assert messages[1]["role"] == "assistant", messages + history_calls: Final = JSON_LIST.validate_python(messages[1]["tool_calls"]) + assert [(call["id"], call["function"]) for call in history_calls] == [ + (CALL_ID, {"name": FLATTENED_NAME, "arguments": ARGUMENTS}) + ], messages + assert messages[2] == {"role": "tool", "tool_call_id": CALL_ID, "content": "5"}, messages + return Reply(body=text_completion(marker)) + + with wire_server(chat_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(model="deepseek/gpt-4o-mini", api_base=wire.url + "/v1") + first: Final = gateway.request( + "POST", + "/v1/responses", + {"model": model, "input": prompt, "tools": [NAMESPACE_TOOL], "store": False}, + ) + assert first.status_code == 200, first.text + first_output: Final = JSON_LIST.validate_python(JSON_OBJECT.validate_json(first.content)["output"]) + calls: Final = tuple(item for item in first_output if item["type"] == "function_call") + assert len(calls) == 1, first.text + assert calls[0]["name"] == TOOL_NAME, first.text + assert calls[0]["namespace"] == NAMESPACE, first.text + assert calls[0]["call_id"] == CALL_ID, first.text + assert calls[0]["arguments"] == ARGUMENTS, first.text + + second: Final = gateway.request( + "POST", + "/v1/responses", + { + "model": model, + "input": [ + {"type": "message", "role": "user", "content": [{"type": "input_text", "text": prompt}]}, + { + "type": "function_call", + "call_id": CALL_ID, + "name": TOOL_NAME, + "namespace": NAMESPACE, + "arguments": ARGUMENTS, + }, + {"type": "function_call_output", "call_id": CALL_ID, "output": "5"}, + ], + "tools": [NAMESPACE_TOOL], + "store": False, + }, + ) + assert second.status_code == 200, second.text + second_output: Final = JSON_LIST.validate_python(JSON_OBJECT.validate_json(second.content)["output"]) + assert [item["type"] for item in second_output] == ["message"], second.text + assert JSON_LIST.validate_python(second_output[0]["content"])[0]["text"] == "The sum is 5", second.text + assert len(wire.drain()) == 2 diff --git a/tests/integration/providers/test_responses_bridge_stream_options.py b/tests/integration/providers/test_responses_bridge_stream_options.py new file mode 100644 index 00000000000..a0efc8d47d0 --- /dev/null +++ b/tests/integration/providers/test_responses_bridge_stream_options.py @@ -0,0 +1,98 @@ +import json +import uuid +from pathlib import Path +from typing import Final + +import pytest +import yaml +from integration._support.client import Gateway +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server + + +def _responses_stream(identity: str, text: str) -> tuple[bytes, ...]: + completed: Final = { + "id": identity, + "object": "response", + "created_at": 1, + "status": "completed", + "model": "gpt-5.3-codex", + "output": [ + { + "type": "message", + "id": f"msg_{identity}", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": text, "annotations": []}], + } + ], + "usage": { + "input_tokens": 11, + "output_tokens": 4, + "total_tokens": 15, + "input_tokens_details": {"cached_tokens": 0}, + "output_tokens_details": {"reasoning_tokens": 0}, + }, + } + events: Final = ( + {"type": "response.created", "response": {**completed, "status": "in_progress", "output": [], "usage": None}}, + { + "type": "response.output_text.delta", + "item_id": f"msg_{identity}", + "output_index": 0, + "content_index": 0, + "delta": text, + }, + {"type": "response.completed", "response": completed}, + ) + return tuple(f"event: {event['type']}\ndata: {json.dumps(event)}\n\n".encode() for event in events) + + +@pytest.mark.covers("providers.responses_bridge.always_include_stream_usage_keeps_include_usage_off_the_responses_wire") +def test_messages_stream_with_always_include_stream_usage_omits_include_usage_from_responses_request( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = "responses-stream-options-" + uuid.uuid4().hex + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/responses", request.target + assert request.headers["authorization"] == "Bearer synthetic-openai-key" + return Reply(content_type="text/event-stream", chunks=_responses_stream(identity, "usage control")) + + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["general_settings"].update({"always_include_stream_usage": True}) + path: Final = tmp_path / "always_include_stream_usage.yaml" + path.write_text(yaml.safe_dump(config)) + with ( + wire_server(respond) as wire, + owned_proxy(gateway, tmp_path, {}, config=path) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(model="openai/gpt-5.3-codex", api_base=wire.url, api_key="synthetic-openai-key") + response: Final = candidate.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 64, + "stream": True, + "messages": [{"role": "user", "content": f"count the usage {identity}"}], + }, + ) + assert response.status_code == 200, response.text + assert "event: message_stop" in response.text, response.text + requests: Final = wire.drain() + assert len(requests) == 1, response.text + assert json.loads(requests[0].body) == { + "model": "gpt-5.3-codex", + "input": [ + { + "type": "message", + "role": "user", + "content": [{"type": "input_text", "text": f"count the usage {identity}"}], + } + ], + "include": ["reasoning.encrypted_content"], + "max_output_tokens": 64, + "stream": True, + }, response.text diff --git a/tests/integration/providers/test_responses_client_header_forwarding_wire.py b/tests/integration/providers/test_responses_client_header_forwarding_wire.py new file mode 100644 index 00000000000..50557cd1727 --- /dev/null +++ b/tests/integration/providers/test_responses_client_header_forwarding_wire.py @@ -0,0 +1,87 @@ +import json +from pathlib import Path +from typing import Final +from uuid import uuid4 + +import pytest +import yaml +from integration._support.client import Gateway +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_BACKEND: Final = "gpt-5.4-mini" +_API_KEY: Final = "synthetic-openai-key" +_CLIENT_HEADER: Final = "x-my-new-header" +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) +_OUTPUT_MESSAGE: Final[dict[str, JsonValue]] = { + "type": "message", + "id": "msg_forwarded", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "header wire control", "annotations": []}], +} +_RESPONSE: Final = json.dumps( + { + "id": "resp_forwarded", + "object": "response", + "status": "completed", + "created_at": 1700000000, + "model": _BACKEND, + "output": [_OUTPUT_MESSAGE], + "usage": { + "input_tokens": 9, + "output_tokens": 3, + "total_tokens": 12, + "input_tokens_details": {"cached_tokens": 0}, + "output_tokens_details": {"reasoning_tokens": 0}, + }, + } +).encode() + + +def _forwarding_config(directory: Path) -> Path: + configuration: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + configuration["general_settings"]["forward_client_headers_to_llm_api"] = True + path: Final = directory / "forwarding.yaml" + path.write_text(yaml.safe_dump(configuration)) + return path + + +@pytest.mark.covers("providers.responses_api.forwarded_client_headers_reach_the_provider") +def test_client_x_header_is_forwarded_to_the_provider_on_responses(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = f"hello-from-client-{uuid4().hex}" + prompt: Final = f"forward my header {marker}" + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/responses", request.target + assert request.headers["authorization"] == f"Bearer {_API_KEY}" + assert request.headers.get(_CLIENT_HEADER) == marker, dict(request.headers) + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["model"] == _BACKEND and body["input"] == prompt, request.body + return Reply(body=_RESPONSE) + + with ( + wire_server(respond) as wire, + owned_proxy(gateway, tmp_path, {}, config=_forwarding_config(tmp_path)) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(model=f"openai/{_BACKEND}", api_base=wire.url, api_key=_API_KEY) + response: Final = candidate.request( + "POST", + "/v1/responses", + {"model": model, "input": prompt, "stream": False}, + headers={_CLIENT_HEADER: marker}, + ) + assert response.status_code == 200, response.text + payload: Final = _JSON_OBJECT.validate_json(response.content) + assert payload["output"] == [ + { + **_OUTPUT_MESSAGE, + "phase": None, + "content": [ + {"type": "output_text", "text": "header wire control", "annotations": [], "logprobs": None} + ], + } + ], response.text + assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/responses")] diff --git a/tests/integration/providers/test_sagemaker_chat_wire.py b/tests/integration/providers/test_sagemaker_chat_wire.py new file mode 100644 index 00000000000..346f4e59e0f --- /dev/null +++ b/tests/integration/providers/test_sagemaker_chat_wire.py @@ -0,0 +1,90 @@ +import json +import uuid +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +_ENDPOINT: Final = "integration-vllm-endpoint" +_INFERENCE_COMPONENT: Final = "integration-vllm-component" +_SERVED_MODEL: Final = "integration-org/served-chat-model" +_ACCESS_KEY: Final = "AKIAINTEGRATION000003" +_PROMPT: Final = "synthetic inference component request" +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + + +def _completion(identity: str) -> bytes: + return json.dumps( + { + "id": identity, + "object": "chat.completion", + "created": 1, + "model": _SERVED_MODEL, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "sagemaker wire control"}, + "finish_reason": "stop", + } + ], + "usage": {"prompt_tokens": 11, "completion_tokens": 4, "total_tokens": 15}, + } + ).encode() + + +@pytest.mark.covers( + "providers.sagemaker_chat_wire.inference_component_header_is_signed_and_hf_model_name_is_the_body_model" +) +def test_sagemaker_chat_signs_the_inference_component_header_and_sends_hf_model_name_as_the_body_model( + gateway: Gateway, +) -> None: + identity: Final = f"sagemaker-chat-{uuid.uuid4().hex}" + + def respond(request: Request) -> Reply: + assert request.method == "POST", request + assert request.target == "/", request + assert request.headers["x-amzn-sagemaker-inference-component"] == _INFERENCE_COMPONENT, dict(request.headers) + authorization: Final = request.headers["authorization"] + assert authorization.startswith(f"AWS4-HMAC-SHA256 Credential={_ACCESS_KEY}/"), authorization + signed_headers: Final = next(part for part in authorization.split(", ") if part.startswith("SignedHeaders=")) + assert "x-amzn-sagemaker-inference-component" in signed_headers.removeprefix("SignedHeaders=").split(";"), ( + authorization + ) + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["model"] == _SERVED_MODEL, body + assert body["messages"] == [{"role": "user", "content": _PROMPT}], body + assert body["max_tokens"] == 16, body + assert "hf_model_name" not in body, body + return Reply(body=_completion(identity)) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=f"sagemaker_chat/{_ENDPOINT}", + api_key=None, + api_base=None, + model_id=_INFERENCE_COMPONENT, + hf_model_name=_SERVED_MODEL, + aws_access_key_id=_ACCESS_KEY, + aws_secret_access_key="synthetic-secret-key-for-testing", + aws_region_name="us-east-1", + sagemaker_base_url=wire.url, + ) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": _PROMPT}], "max_tokens": 16}, + ) + assert response.status_code == 200, response.text + payload: Final = _JSON_OBJECT.validate_json(response.content) + assert payload["id"] == identity, response.text + assert payload["choices"] == [ + { + "finish_reason": "stop", + "index": 0, + "message": {"role": "assistant", "content": "sagemaker wire control"}, + "provider_specific_fields": {}, + } + ], response.text + assert [(request.method, request.target) for request in wire.drain()] == [("POST", "/")], response.text diff --git a/tests/integration/providers/test_vertex_batch_output_info_wire.py b/tests/integration/providers/test_vertex_batch_output_info_wire.py new file mode 100644 index 00000000000..a7ac896076c --- /dev/null +++ b/tests/integration/providers/test_vertex_batch_output_info_wire.py @@ -0,0 +1,124 @@ +import base64 +import functools +import json +from typing import Final + +import pytest +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server + +PROJECT: Final = "cc-scripted-project" +LOCATION: Final = "us-central1" +MODEL: Final = "vertex_ai/gemini-2.5-flash" +VERTEX_MODEL_RESOURCE: Final = "publishers/google/models/gemini-2.5-flash" +BUCKET: Final = "integration-batch-bucket" +INPUT_FILE_ID: Final = f"gs://{BUCKET}/litellm-vertex-files/{VERTEX_MODEL_RESOURCE}/input.jsonl" +OUTPUT_PREFIX: Final = INPUT_FILE_ID.rsplit("/", 1)[0] +JOB_NAME: Final = f"projects/{PROJECT}/locations/{LOCATION}/batchPredictionJobs/7412345678901234567" +JOB_ID: Final = JOB_NAME.rsplit("/", 1)[-1] +EXPECTED_VERTEX_BODY: Final = { + "inputConfig": {"gcsSource": {"uris": [INPUT_FILE_ID]}, "instancesFormat": "jsonl"}, + "outputConfig": {"predictionsFormat": "jsonl", "gcsDestination": {"outputUriPrefix": OUTPUT_PREFIX}}, + "model": VERTEX_MODEL_RESOURCE, +} +VERTEX_REPLY: Final = { + "name": JOB_NAME, + "displayName": "litellm-vertex-batch-scripted", + "model": VERTEX_MODEL_RESOURCE, + "inputConfig": {"gcsSource": {"uris": [INPUT_FILE_ID]}, "instancesFormat": "jsonl"}, + "outputConfig": {"predictionsFormat": "jsonl", "gcsDestination": {"outputUriPrefix": OUTPUT_PREFIX}}, + "outputInfo": None, + "state": "JOB_STATE_PENDING", + "createTime": "2026-07-24T20:00:00.000000Z", + "updateTime": "2026-07-24T20:00:00.000000Z", +} + + +@functools.cache +def _vertex_private_key_pem() -> str: + return ( + rsa.generate_private_key(public_exponent=65537, key_size=2048) + .private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ) + .decode() + ) + + +def _vertex_service_account_json(url: str) -> str: + return json.dumps( + { + "type": "service_account", + "project_id": PROJECT, + "private_key_id": "scripted", + "private_key": _vertex_private_key_pem(), + "client_email": f"scripted@{PROJECT}.iam.gserviceaccount.com", + "client_id": "0", + "auth_uri": f"{url}/_oauth/authorize", + "token_uri": f"{url}/_oauth/token", + } + ) + + +def _encoded(raw: str, model: str, prefix: str) -> str: + return prefix + base64.urlsafe_b64encode(f"litellm:{raw};model,{model}".encode()).decode().rstrip("=") + + +def vertex_peer(request: Request) -> Reply: + assert request.method == "POST", request.method + assert request.target == f"/v1/projects/{PROJECT}/locations/{LOCATION}/batchPredictionJobs", request.target + assert request.headers["authorization"] == "Bearer scripted-token" + assert request.headers["content-type"] == "application/json; charset=utf-8" + body: Final = json.loads(request.body) + display_name: Final = body.pop("displayName") + assert isinstance(display_name, str) and display_name.startswith("litellm-vertex-batch-"), display_name + assert body == EXPECTED_VERTEX_BODY, body + return Reply(body=json.dumps(VERTEX_REPLY).encode()) + + +@pytest.mark.covers("other.provider_wire.vertex_ai.batch_create_with_null_output_info_returns_batch_instead_of_500") +def test_vertex_batch_create_survives_explicit_null_output_info(gateway: Gateway) -> None: + with wire_server(vertex_peer) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=MODEL, + api_key=None, + api_base=wire.url, + vertex_project=PROJECT, + vertex_location=LOCATION, + vertex_credentials=_vertex_service_account_json(gateway.upstream_url), + ) + response: Final = gateway.request( + "POST", + "/v1/batches", + { + "input_file_id": INPUT_FILE_ID, + "endpoint": "/v1/chat/completions", + "completion_window": "24h", + "model": model, + }, + ) + assert response.status_code == 200, response.text + body: Final = response.json() + assert ( + body["id"], + body["object"], + body["status"], + body["input_file_id"], + body["output_file_id"], + body["error_file_id"], + body["completion_window"], + ) == ( + _encoded(JOB_ID, model, "batch_"), + "batch", + "validating", + _encoded(INPUT_FILE_ID, model, "file-"), + _encoded(f"{OUTPUT_PREFIX}/predictions.jsonl", model, "file-"), + None, + "24h", + ), response.text + requests: Final = wire.drain() + assert len(requests) == 1, f"Expected exactly one Vertex POST, saw {[request.target for request in requests]}" diff --git a/tests/integration/providers/test_vertex_gemini_fragmented_stream_wire.py b/tests/integration/providers/test_vertex_gemini_fragmented_stream_wire.py new file mode 100644 index 00000000000..449c5c9c105 --- /dev/null +++ b/tests/integration/providers/test_vertex_gemini_fragmented_stream_wire.py @@ -0,0 +1,138 @@ +import json +import time +from typing import Final + +import pytest +from cryptography.hazmat.primitives import serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server +from pydantic import BaseModel, ConfigDict, JsonValue, TypeAdapter + +_BACKEND: Final = "gemini-3.7-flash" +_PROJECT: Final = "scripted-project" +_LOCATION: Final = "us-central1" +_MODEL_PATH: Final = f"/v1/projects/{_PROJECT}/locations/{_LOCATION}/publishers/google/models/{_BACKEND}" +_PROMPT: Final = "Write a very long numbered list." +_PART_COUNT: Final = 8000 +_LINES_PER_FRAGMENT: Final = 64 +_STREAM_BUDGET_SECONDS: Final = 10.0 +_JSON_OBJECT: Final = TypeAdapter(dict[str, JsonValue]) + + +class _Delta(BaseModel): + model_config = ConfigDict(extra="ignore") + content: str | None = None + + +class _Choice(BaseModel): + model_config = ConfigDict(extra="ignore") + delta: _Delta + finish_reason: str | None = None + + +class _Chunk(BaseModel): + model_config = ConfigDict(extra="ignore") + choices: tuple[_Choice, ...] + + +def _service_account_json(token_url: str) -> str: + private_key: Final = ( + rsa.generate_private_key(public_exponent=65537, key_size=2048) + .private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.PKCS8, + serialization.NoEncryption(), + ) + .decode() + ) + return json.dumps( + { + "type": "service_account", + "project_id": _PROJECT, + "private_key_id": "scripted", + "private_key": private_key, + "client_email": f"scripted@{_PROJECT}.iam.gserviceaccount.com", + "client_id": "0", + "auth_uri": f"{token_url}/_oauth/authorize", + "token_uri": f"{token_url}/_oauth/token", + } + ) + + +def _expected_text() -> str: + return "".join(f"{index}. item\n" for index in range(_PART_COUNT)) + + +def _gemini_response_fragments() -> tuple[bytes, ...]: + document: Final = json.dumps( + { + "candidates": [ + { + "content": { + "role": "model", + "parts": [{"text": f"{index}. item\n"} for index in range(_PART_COUNT)], + }, + "finishReason": "STOP", + } + ], + "usageMetadata": {"promptTokenCount": 9, "candidatesTokenCount": 40000, "totalTokenCount": 40009}, + "modelVersion": _BACKEND, + }, + indent=2, + ) + lines: Final = document.split("\n") + fragments: Final = tuple( + "\n".join(lines[start : start + _LINES_PER_FRAGMENT]).encode() + b"\n" + for start in range(0, len(lines), _LINES_PER_FRAGMENT) + ) + return (b"data: " + fragments[0], *fragments[1:], b"\n") + + +@pytest.mark.covers("providers.vertex_gemini.fragmented_stream_json_is_parsed_once_and_stays_live") +def test_vertex_gemini_stream_split_across_many_fragments_completes_without_stalling(gateway: Gateway) -> None: + fragments: Final = _gemini_response_fragments() + + def respond(request: Request) -> Reply: + assert request.method == "POST" + assert request.target == f"{_MODEL_PATH}:streamGenerateContent?alt=sse" + assert request.headers["authorization"] == "Bearer scripted-token" + body: Final = _JSON_OBJECT.validate_json(request.body) + assert body["contents"] == [{"role": "user", "parts": [{"text": _PROMPT}]}] + assert body["generationConfig"] == {"temperature": 0.0} + return Reply(content_type="text/event-stream", chunks=fragments) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=f"vertex_ai/{_BACKEND}", + api_base=f"{wire.url}{_MODEL_PATH}", + api_key=None, + vertex_project=_PROJECT, + vertex_location=_LOCATION, + vertex_credentials=_service_account_json(gateway.upstream_url.rstrip("/")), + ) + started: Final = time.monotonic() + with gateway.client.stream( + "POST", + "/v1/chat/completions", + json={ + "model": model, + "messages": [{"role": "user", "content": _PROMPT}], + "stream": True, + "temperature": 0.0, + }, + headers={"Authorization": f"Bearer {gateway.key}"}, + timeout=_STREAM_BUDGET_SECONDS, + ) as response: + assert response.status_code == 200, response.read() + lines: Final = tuple(line for line in response.iter_lines() if line.startswith("data: ")) + elapsed: Final = time.monotonic() - started + assert elapsed < _STREAM_BUDGET_SECONDS, f"stream took {elapsed:.1f}s for {len(fragments)} fragments" + assert lines[-1] == "data: [DONE]", lines[-3:] + chunks: Final = tuple(_Chunk.model_validate_json(line.removeprefix("data: ")) for line in lines[:-1]) + choices: Final = tuple(choice for chunk in chunks for choice in chunk.choices) + assert "".join(choice.delta.content or "" for choice in choices) == _expected_text() + assert tuple(choice.finish_reason for choice in choices if choice.finish_reason) == ("stop",) + assert [(request.method, request.target) for request in wire.drain()] == [ + ("POST", f"{_MODEL_PATH}:streamGenerateContent?alt=sse") + ] diff --git a/tests/integration/providers/test_websearch_interception_wire.py b/tests/integration/providers/test_websearch_interception_wire.py index cba4f174233..a6f098cf64f 100644 --- a/tests/integration/providers/test_websearch_interception_wire.py +++ b/tests/integration/providers/test_websearch_interception_wire.py @@ -183,6 +183,8 @@ from integration._support.client import Gateway, eventually _QUERY: Final = "integration capped search" _TEXT_BLOCK: Final = {"type": "text", "text": "searching once more"} _NOT_INTERCEPTED: Final = "native tool reached the provider" +_FINAL_BLOCK: Final = {"type": "text", "text": "answered from the stored backend"} +_OWNED_RESULT_TEXT: Final = "Title: Owned result\nURL: https://owned.invalid/a\nSnippet: owned snippet" _SEARCH_RESULT_BLOCK: Final = { "type": "web_search_result", "url": "https://owned.invalid/a", @@ -297,3 +299,111 @@ def test_capped_websearch_interception_loop_ends_turn_instead_of_exposing_intern assert content[2] == _TEXT_BLOCK, response.text targets: Final = tuple((request.method, urlsplit(request.target).path) for request in wire.drain()) assert targets[-3:] == (("POST", "/v1/messages"), ("GET", "/search"), ("POST", "/v1/messages")), targets + + +@pytest.mark.covers("other.provider_wire.anthropic.websearch_interception_uses_database_search_tool_backend") +def test_database_created_search_tool_backend_receives_the_intercepted_query_over_a_same_named_config_tool( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = "websearch-db-" + uuid.uuid4().hex + tool_name: Final = "integration-db-searxng-" + uuid.uuid4().hex + searched: Final = threading.Event() + + def respond(request: Request) -> Reply: + parts: Final = urlsplit(request.target) + if request.method == "GET" and parts.path == "/database/search": + assert parse_qs(parts.query)["q"] == [_QUERY], request.target + searched.set() + return Reply( + body=json.dumps( + { + "results": [ + {"title": "Owned result", "url": "https://owned.invalid/a", "content": "owned snippet"} + ] + } + ).encode() + ) + assert request.method == "POST" and parts.path == "/v1/messages", request.target + body: Final = json.loads(request.body) + if any(tool.get("type") == "web_search_20250305" for tool in body["tools"]): + return _anthropic_reply(identity, [{"type": "text", "text": _NOT_INTERCEPTED}], "end_turn") + assert [tool["name"] for tool in body["tools"]] == ["litellm_web_search"], body["tools"] + results: Final = [ + block + for message in body["messages"] + if isinstance(message["content"], list) + for block in message["content"] + if block["type"] == "tool_result" + ] + if not results: + return _anthropic_reply(identity, [_TEXT_BLOCK, _search_tool_use(identity)], "tool_use") + assert results == [{"type": "tool_result", "tool_use_id": identity, "content": _OWNED_RESULT_TEXT}], results + return _anthropic_reply(identity, [_FINAL_BLOCK], "end_turn") + + def send(candidate: Gateway, model: str) -> httpx.Response: + return candidate.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 64, + "messages": [{"role": "user", "content": identity + " attempt " + uuid.uuid4().hex}], + "tools": [{"type": "web_search_20250305", "name": "web_search", "max_uses": 3}], + }, + ) + + def searched_through_proxy(response: httpx.Response) -> bool: + return searched.is_set() and _NOT_INTERCEPTED not in response.text + + with wire_server(respond) as wire, gateway.scenario() as scenario: + created: Final = gateway.post( + "/search_tools", + { + "search_tool": { + "search_tool_name": tool_name, + "litellm_params": {"search_provider": "searxng", "api_base": wire.url + "/database"}, + } + }, + ) + scenario.cleanups.callback(gateway.request, "DELETE", f"/search_tools/{created['search_tool_id']}") + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["search_tools"] = [ + { + "search_tool_name": tool_name, + "litellm_params": {"search_provider": "searxng", "api_base": wire.url + "/config"}, + } + ] + config["litellm_settings"].update( + { + "callbacks": ["websearch_interception"], + "websearch_interception_params": { + "enabled": True, + "enabled_providers": ["anthropic"], + "search_tool_name": tool_name, + }, + } + ) + path: Final = tmp_path / "websearch-db.yaml" + path.write_text(yaml.safe_dump(config)) + environment: Final = {"ANTHROPIC_API_BASE": wire.url} + with owned_proxy(gateway, tmp_path, environment, config=path) as candidate, candidate.scenario() as models: + model: Final = models.model( + model="anthropic/claude-sonnet-4-5-20250929", api_base=wire.url, api_key="synthetic-anthropic-key" + ) + response: Final = eventually(lambda: send(candidate, model), searched_through_proxy, seconds=40) + assert response.status_code == 200, response.text + body: Final = response.json() + assert body["stop_reason"] == "end_turn", response.text + assert body["content"][-1] == _FINAL_BLOCK, response.text + found: Final = [ + (result["url"], result["title"]) + for block in body["content"] + if block["type"] == "web_search_tool_result" + for result in block["content"] + ] + assert found == [("https://owned.invalid/a", "Owned result")], response.text + assert "litellm_web_search" not in response.text, response.text + targets: Final = tuple((request.method, urlsplit(request.target).path) for request in wire.drain()) + assert targets[-3:] == (("POST", "/v1/messages"), ("GET", "/database/search"), ("POST", "/v1/messages")), ( + targets + ) diff --git a/tests/integration/routing/test_advisor_failure_cooldown.py b/tests/integration/routing/test_advisor_failure_cooldown.py new file mode 100644 index 00000000000..41618e29109 --- /dev/null +++ b/tests/integration/routing/test_advisor_failure_cooldown.py @@ -0,0 +1,101 @@ +import json +import uuid +from pathlib import Path +from typing import Final + +import pytest +import yaml +from integration._support.client import Gateway +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server + +_ADVISOR_KEY: Final = "synthetic-advisor-key" +_QUESTION: Final = "which index should this query use" +_PROXY_CONFIG: Final = Path(__file__).resolve().parents[1] / "proxy_config.yaml" + + +def _executor_reply(body: dict[str, object], identity: str) -> Reply: + tools: Final = body.get("tools") + message: Final = ( + { + "role": "assistant", + "content": None, + "tool_calls": [ + { + "id": "advisor-call", + "type": "function", + "function": {"name": "advisor", "arguments": json.dumps({"question": _QUESTION})}, + } + ], + } + if isinstance(tools, list) + else {"role": "assistant", "content": "served without an advisor"} + ) + return Reply( + body=json.dumps( + { + "id": f"chatcmpl-{identity}-{uuid.uuid4().hex[:8]}", + "object": "chat.completion", + "created": 1, + "model": "llama-3.3-70b-versatile", + "choices": [{"index": 0, "message": message, "finish_reason": "tool_calls" if tools else "stop"}], + "usage": {"prompt_tokens": 10, "completion_tokens": 4, "total_tokens": 14}, + } + ).encode() + ) + + +def _cooldowns_enabled_config(directory: Path) -> Path: + loaded: Final = yaml.safe_load(_PROXY_CONFIG.read_text()) + path: Final = directory / "cooldowns_enabled.yaml" + path.write_text(yaml.safe_dump({**loaded, "router_settings": {"num_retries": 0}})) + return path + + +@pytest.mark.covers("routing.cooldown.advisor_sub_call_failure_does_not_cool_down_the_executor_deployment") +def test_advisor_sub_call_401_leaves_the_executor_deployment_serving_the_next_request( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = "advisor-cooldown-" + uuid.uuid4().hex + + def respond(request: Request) -> Reply: + if request.target == "/v1/chat/completions": + return _executor_reply(json.loads(request.body), identity) + assert request.target == "/v1/messages" + assert request.headers["x-api-key"] == _ADVISOR_KEY + return Reply( + status=401, + body=json.dumps( + {"type": "error", "error": {"type": "authentication_error", "message": "invalid x-api-key"}} + ).encode(), + ) + + with ( + wire_server(respond) as wire, + owned_proxy(gateway, tmp_path, {}, config=_cooldowns_enabled_config(tmp_path)) as candidate, + candidate.scenario() as scenario, + ): + executor: Final = scenario.model(model="hosted_vllm/gpt-4o-mini", api_base=wire.url + "/v1") + advisor: Final = scenario.model( + model="anthropic/claude-opus-4-1-20250805", api_base=wire.url, api_key=_ADVISOR_KEY + ) + advised: Final = candidate.request( + "POST", + "/v1/messages", + { + "model": executor, + "max_tokens": 64, + "messages": [{"role": "user", "content": identity}], + "tools": [{"type": "advisor_20260301", "name": "advisor", "model": advisor}], + }, + ) + assert advised.status_code == 401, advised.text + assert [request.target for request in wire.drain()] == ["/v1/chat/completions", "/v1/messages"] + unrelated: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": executor, "messages": [{"role": "user", "content": identity + " unrelated"}]}, + ) + assert unrelated.status_code == 200, unrelated.text + assert unrelated.json()["choices"][0]["message"]["content"] == "served without an advisor", unrelated.text + assert [request.target for request in wire.drain()] == ["/v1/chat/completions"] diff --git a/tests/integration/routing/test_key_tpm_reservation.py b/tests/integration/routing/test_key_tpm_reservation.py new file mode 100644 index 00000000000..8d0e06679cd --- /dev/null +++ b/tests/integration/routing/test_key_tpm_reservation.py @@ -0,0 +1,59 @@ +import json +import time +import uuid +from collections import Counter +from concurrent.futures import ThreadPoolExecutor +from typing import Final + +import httpx +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue + +KEY_TPM_LIMIT: Final = 100 +MAX_TOKENS: Final = 80 +CONCURRENT_REQUESTS: Final = 10 +PROVIDER_HOLD_SECONDS: Final = 2.0 +UPSTREAM_REPLY: Final = json.dumps( + { + "id": "chatcmpl_tpm_reservation", + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [{"index": 0, "message": {"role": "assistant", "content": "reserved"}, "finish_reason": "stop"}], + "usage": {"prompt_tokens": 20, "completion_tokens": 20, "total_tokens": 40}, + } +).encode() + + +@pytest.mark.covers("quota_management.key_tpm_limit.concurrent_requests_reserve_tokens_before_provider_call") +def test_concurrent_requests_over_key_tpm_are_rejected_before_reaching_provider(gateway: Gateway) -> None: + probe: Final = "tpm reservation probe " + uuid.uuid4().hex[:8] + messages: Final[list[JsonValue]] = [{"role": "user", "content": probe}] + + def respond(request: Request) -> Reply: + assert (request.method, request.target) == ("POST", "/v1/chat/completions") + assert json.loads(request.body) == {"model": "gpt-4o-mini", "max_tokens": MAX_TOKENS, "messages": messages} + time.sleep(PROVIDER_HOLD_SECONDS) + return Reply(body=UPSTREAM_REPLY) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model(api_base=f"{wire.url}/v1") + key: Final = scenario.key(tpm_limit=KEY_TPM_LIMIT) + body: Final[dict[str, JsonValue]] = { + "model": model, + "max_tokens": MAX_TOKENS, + "messages": messages, + } + + def send(_: int) -> httpx.Response: + return gateway.request("POST", "/v1/chat/completions", body, key=key) + + with ThreadPoolExecutor(max_workers=CONCURRENT_REQUESTS) as pool: + responses: Final = tuple(pool.map(send, range(CONCURRENT_REQUESTS))) + statuses: Final = Counter(response.status_code for response in responses) + assert statuses == Counter({200: 1, 429: CONCURRENT_REQUESTS - 1}), tuple( + response.text for response in responses + ) + assert tuple(json.loads(request.body)["messages"] for request in wire.drain()) == (messages,) diff --git a/tests/integration/routing/test_priority_model_tpm_enforcement.py b/tests/integration/routing/test_priority_model_tpm_enforcement.py new file mode 100644 index 00000000000..c3d0446c1e1 --- /dev/null +++ b/tests/integration/routing/test_priority_model_tpm_enforcement.py @@ -0,0 +1,116 @@ +import json +import uuid +from pathlib import Path +from queue import SimpleQueue +from typing import Final + +import httpx +import pytest +import yaml +from integration._support.client import Gateway, eventually +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server + +OPENAI_MODEL: Final = "gpt-4.1-mini" +PROMPT_TOKENS: Final = 30 +COMPLETION_TOKENS: Final = 10 +MODEL_TPM: Final = PROMPT_TOKENS + COMPLETION_TOKENS +PREMIUM_SHARE: Final = 0.5 +UPSTREAM_REPLY: Final = json.dumps( + { + "id": "chatcmpl_model_tpm_enforcement", + "object": "chat.completion", + "created": 1700000000, + "model": OPENAI_MODEL, + "choices": [ + { + "index": 0, + "message": {"role": "assistant", "content": "model tpm control"}, + "finish_reason": "stop", + } + ], + "usage": { + "prompt_tokens": PROMPT_TOKENS, + "completion_tokens": COMPLETION_TOKENS, + "total_tokens": PROMPT_TOKENS + COMPLETION_TOKENS, + }, + } +).encode() + + +@pytest.mark.covers("other.routing.priority_rate_limits.tpm_only_model_rejects_priority_traffic_at_capacity") +def test_tpm_only_model_returns_429_to_priority_key_once_recorded_tokens_reach_model_tpm( + gateway: Gateway, tmp_path: Path +) -> None: + probe: Final = "model tpm probe " + uuid.uuid4().hex + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/chat/completions" + assert request.headers["authorization"] == "Bearer synthetic-openai-key" + body: Final = json.loads(request.body) + assert body["messages"][0]["content"].startswith(probe), body + assert body == { + "model": OPENAI_MODEL, + "messages": [{"role": "user", "content": body["messages"][0]["content"]}], + "max_tokens": 16, + } + return Reply(body=UPSTREAM_REPLY) + + configuration: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + configuration["litellm_settings"] = { + **configuration["litellm_settings"], + "callbacks": ["dynamic_rate_limiter_v3"], + "priority_reservation": {"premium": PREMIUM_SHARE}, + } + path: Final = tmp_path / "priority.yaml" + path.write_text(yaml.safe_dump(configuration)) + with ( + wire_server(respond) as wire, + owned_proxy(gateway, tmp_path, {}, config=path) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model( + model=f"openai/{OPENAI_MODEL}", + api_base=f"{wire.url}/v1", + api_key="synthetic-openai-key", + tpm=MODEL_TPM, + ) + key: Final = scenario.key(metadata={"priority": "premium"}) + responses: Final[SimpleQueue[httpx.Response]] = SimpleQueue() + + def attempt() -> httpx.Response: + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "max_tokens": 16, + "messages": [{"role": "user", "content": f"{probe} {uuid.uuid4().hex}"}], + }, + key=key, + ) + responses.put(response) + return response + + first: Final = attempt() + assert first.status_code == 200, first.text + assert first.json()["usage"]["total_tokens"] == MODEL_TPM, first.text + blocked: Final = eventually(attempt, lambda response: response.status_code == 429, seconds=30) + served: Final = tuple(responses.get_nowait() for _ in range(responses.qsize())) + assert all(response.status_code == 200 for response in served[:-1]), [r.status_code for r in served] + assert len(wire.drain()) == len(served) - 1 + assert blocked.headers["x-litellm-priority"] == "premium", blocked.headers + assert blocked.headers["rate_limit_type"] == "tokens", blocked.headers + detail: Final = ( + f"Model capacity reached for {model}. Priority: premium, Rate limit type: tokens, " + f"Model TPM: {MODEL_TPM}, Model RPM: not configured, Remaining: 0" + ) + assert blocked.json() == { + "error": { + "message": detail, + "type": "throttling_error", + "param": None, + "code": "429", + "provider_specific_fields": {"error": detail}, + } + }, blocked.text diff --git a/tests/integration/routing/test_priority_rate_limit_headers.py b/tests/integration/routing/test_priority_rate_limit_headers.py index 2d21f3ba8b2..bd92a362885 100644 --- a/tests/integration/routing/test_priority_rate_limit_headers.py +++ b/tests/integration/routing/test_priority_rate_limit_headers.py @@ -5,7 +5,7 @@ from typing import Final import pytest import yaml -from integration._support.client import Gateway +from integration._support.client import Gateway, eventually from integration._support.process import owned_proxy from integration._support.wire import Reply, Request, wire_server @@ -27,6 +27,27 @@ UPSTREAM_REPLY: Final = json.dumps( ).encode() +CHAT_MODEL: Final = "gpt-5.6" +MAX_COMPLETION_TOKENS: Final = 64 + + +def _chat_frames(identity: str, text: str) -> tuple[bytes, ...]: + events: Final = ( + {"choices": [{"index": 0, "delta": {"role": "assistant", "content": text}, "finish_reason": None}]}, + {"choices": [{"index": 0, "delta": {}, "finish_reason": "stop"}]}, + {"choices": [], "usage": {"prompt_tokens": 10, "completion_tokens": 4, "total_tokens": 14}}, + ) + frames: Final = tuple( + b"data: " + + json.dumps( + {"id": identity, "object": "chat.completion.chunk", "created": 1, "model": CHAT_MODEL, **event} + ).encode() + + b"\n\n" + for event in events + ) + return (*frames, b"data: [DONE]\n\n") + + @pytest.mark.covers("other.routing.priority_rate_limits.v1_messages_success_exposes_v3_priority_headers") def test_non_streaming_v1_messages_success_carries_v3_priority_rate_limit_headers( gateway: Gateway, tmp_path: Path @@ -86,3 +107,89 @@ def test_non_streaming_v1_messages_success_carries_v3_priority_rate_limit_header } observed: Final = {name: response.headers.get(name) for name in expected} assert observed == expected, response.headers + + +@pytest.mark.covers("other.routing.priority_rate_limits.streaming_success_logs_v3_remaining_values_for_callbacks") +def test_streaming_chat_completion_success_logs_v3_rate_limit_remaining_values_for_callbacks( + gateway: Gateway, tmp_path: Path +) -> None: + probe: Final = "streaming remaining probe " + uuid.uuid4().hex + sink_secret: Final = "synthetic-sink-secret-" + uuid.uuid4().hex + + def provider(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/chat/completions", request.target + assert request.headers["authorization"] == "Bearer synthetic-openai-key" + assert json.loads(request.body) == { + "model": CHAT_MODEL, + "messages": [{"role": "user", "content": probe}], + "max_completion_tokens": MAX_COMPLETION_TOKENS, + "stream": True, + "stream_options": {"include_usage": True}, + }, request.body + return Reply(content_type="text/event-stream", chunks=_chat_frames("chatcmpl_" + probe[-8:], "streamed")) + + def sink(request: Request) -> Reply: + assert request.headers["authorization"] == f"Bearer {sink_secret}" + return Reply() + + configuration: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + configuration["litellm_settings"] = { + **configuration["litellm_settings"], + "callbacks": ["generic_api"], + "DEFAULT_FLUSH_INTERVAL_SECONDS": 1, + } + path: Final = tmp_path / "per_key_streaming.yaml" + path.write_text(yaml.safe_dump(configuration)) + with ( + wire_server(provider) as wire, + wire_server(sink) as endpoint, + owned_proxy( + gateway, + tmp_path, + {"GENERIC_LOGGER_ENDPOINT": endpoint.url, "GENERIC_LOGGER_HEADERS": f"Authorization=Bearer {sink_secret}"}, + config=path, + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model( + model=f"openai/{CHAT_MODEL}", + api_base=wire.url + "/v1", + api_key="synthetic-openai-key", + ) + key: Final = scenario.key(model_rpm_limit={model: MODEL_RPM}, model_tpm_limit={model: MODEL_TPM}) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "messages": [{"role": "user", "content": probe}], + "max_completion_tokens": MAX_COMPLETION_TOKENS, + "stream": True, + "stream_options": {"include_usage": True}, + }, + key=key, + ) + assert response.status_code == 200, response.text + assert '"content":"streamed"' in response.text, response.text + assert len(wire.drain()) == 1 + batches: Final[ + list[Request] + ] = [] # mutable-ok: drain() consumes the queue, later polls must keep earlier batches + + def delivered() -> tuple[dict, ...]: + batches.extend(endpoint.drain()) + return tuple( + event for batch in batches for event in json.loads(batch.body) if event.get("model_group") == model + ) + + events: Final = eventually(delivered, lambda values: len(values) == 1, seconds=10) + assert (events[0]["status"], events[0]["stream"]) == ("success", True), json.dumps(events[0]) + additional_headers: Final = events[0]["hidden_params"]["additional_headers"] or {} + observed: Final = {name: value for name, value in additional_headers.items() if name.startswith("x-ratelimit-")} + remaining_tokens: Final = observed.get("x-ratelimit-model_per_key-remaining-tokens") + assert isinstance(remaining_tokens, int) and 0 < remaining_tokens <= MODEL_TPM, json.dumps(observed) + assert {name: value for name, value in observed.items() if not name.endswith("-remaining-tokens")} == { + "x-ratelimit-model_per_key-limit-requests": MODEL_RPM, + "x-ratelimit-model_per_key-remaining-requests": MODEL_RPM - 1, + "x-ratelimit-model_per_key-limit-tokens": MODEL_TPM, + }, json.dumps(events[0]["hidden_params"]) diff --git a/tests/integration/routing/test_team_model_tpm_limit.py b/tests/integration/routing/test_team_model_tpm_limit.py new file mode 100644 index 00000000000..741c41c9285 --- /dev/null +++ b/tests/integration/routing/test_team_model_tpm_limit.py @@ -0,0 +1,89 @@ +import json +import threading +import uuid +from concurrent.futures import Future, ThreadPoolExecutor +from typing import Final + +import httpx +import pytest +from integration._support.client import Gateway, eventually +from integration._support.wire import Reply, Request, wire_server + +PROVIDER_MODEL: Final = "gpt-4o-mini" +TEAM_MODEL_TPM: Final = 100 +MAX_TOKENS: Final = 60 +CONCURRENT_REQUESTS: Final = 3 +UPSTREAM_REPLY: Final = json.dumps( + { + "id": "chatcmpl-team-tpm-control", + "object": "chat.completion", + "created": 1, + "model": PROVIDER_MODEL, + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "team tpm control"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": 10, "completion_tokens": 4, "total_tokens": 14}, + } +).encode() + + +@pytest.mark.covers("routing.team_model_tpm.concurrent_requests_over_the_limit_are_rejected_before_the_provider_call") +def test_concurrent_team_model_tpm_requests_reserve_tokens_before_reaching_the_provider(gateway: Gateway) -> None: + probe: Final = "team tpm probe " + uuid.uuid4().hex + release: Final = threading.Event() + + def respond(request: Request) -> Reply: + assert request.method == "POST" and request.target == "/v1/chat/completions" + assert request.headers["authorization"] == "Bearer synthetic-team-tpm-key" + body: Final = json.loads(request.body) + content: Final = body["messages"][0]["content"] + assert body == { + "model": PROVIDER_MODEL, + "messages": [{"role": "user", "content": content}], + "max_tokens": MAX_TOKENS, + } + assert content.startswith(probe), content + release.wait(timeout=10) + return Reply(body=UPSTREAM_REPLY) + + with wire_server(respond) as wire, gateway.scenario() as scenario: + model: Final = scenario.model( + model=f"openai/{PROVIDER_MODEL}", + api_base=wire.url + "/v1", + api_key="synthetic-team-tpm-key", + ) + team: Final = scenario.team(metadata={"model_tpm_limit": {model: TEAM_MODEL_TPM}}) + key: Final = scenario.key(team_id=team) + + def send(index: int) -> httpx.Response: + return gateway.request( + "POST", + "/v1/chat/completions", + { + "model": model, + "max_tokens": MAX_TOKENS, + "messages": [{"role": "user", "content": f"{probe} {index}"}], + }, + key=key, + ) + + with ThreadPoolExecutor(max_workers=CONCURRENT_REQUESTS) as pool: + futures: Final[tuple[Future[httpx.Response], ...]] = tuple( + pool.submit(send, index) for index in range(CONCURRENT_REQUESTS) + ) + eventually( + lambda: sum(future.done() for future in futures) + wire.received.qsize(), + lambda settled: settled >= CONCURRENT_REQUESTS, + seconds=10, + ) + release.set() + responses: Final = tuple(future.result(timeout=15) for future in futures) + statuses: Final = tuple(sorted(response.status_code for response in responses)) + assert statuses == (200, 429, 429), tuple(response.text for response in responses) + assert len(wire.drain()) == 1, statuses + served: Final = next(response for response in responses if response.status_code == 200) + assert served.json()["usage"] == {"prompt_tokens": 10, "completion_tokens": 4, "total_tokens": 14}, served.text + for rejected in (response for response in responses if response.status_code == 429): + error: Final = rejected.json()["error"] + assert (error["type"], error["code"], error["param"]) == ("throttling_error", "429", None), rejected.text + assert f"Limit type: tokens. Current limit: {TEAM_MODEL_TPM}," in error["message"], rejected.text diff --git a/tests/integration/sdk/test_aiohttp_session_rebuild_wire.py b/tests/integration/sdk/test_aiohttp_session_rebuild_wire.py new file mode 100644 index 00000000000..1ec1f486bed --- /dev/null +++ b/tests/integration/sdk/test_aiohttp_session_rebuild_wire.py @@ -0,0 +1,109 @@ +from __future__ import annotations + +import json +import os +import subprocess +import sys +import textwrap +import threading +from collections.abc import Iterator +from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer +from typing import Final + +import pytest +from pydantic import JsonValue, TypeAdapter + +CONFIGURED_KEEPALIVE_SECONDS: Final = 1 +IDLE_SECONDS: Final = 2 +RESPONSES: Final = TypeAdapter(list[dict[str, JsonValue]]) + +REBUILT_SESSION_EXCHANGE: Final = textwrap.dedent( + """ + import asyncio, json, sys + from aiohttp import ClientSession + from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler + + async def main(base_url: str, idle_seconds: float) -> None: + shared = ClientSession() + handler = AsyncHTTPHandler(shared_session=shared) + await shared.close() + first = await handler.post(f"{base_url}/embeddings", json={"input": "warm-up"}) + await asyncio.sleep(idle_seconds) + second = await handler.post(f"{base_url}/embeddings", json={"input": "warm-up"}) + print(json.dumps([first.json(), second.json()])) + await handler.close() + + asyncio.run(main(sys.argv[1], float(sys.argv[2]))) + """ +) + + +class _ConnectionCountingPeer(ThreadingHTTPServer): + daemon_threads = True + + def __init__(self, address: tuple[str, int]) -> None: + super().__init__(address, _ConnectionHandler) + self.lock = threading.Lock() + self.connections = 0 + + def next_connection(self) -> int: + with self.lock: + self.connections += 1 + return self.connections + + +class _ConnectionHandler(BaseHTTPRequestHandler): + protocol_version = "HTTP/1.1" + server: _ConnectionCountingPeer + + def setup(self) -> None: + super().setup() + self.connection_number = self.server.next_connection() + + def do_POST(self) -> None: + self.rfile.read(int(self.headers["Content-Length"])) + body: Final = json.dumps({"connection": self.connection_number}).encode() + self.send_response(200) + self.send_header("Content-Type", "application/json") + self.send_header("Content-Length", str(len(body))) + self.end_headers() + self.wfile.write(body) + + def log_message(self, format: str, *args: object) -> None: + return + + +@pytest.fixture +def connection_counting_peer() -> Iterator[str]: + server: Final = _ConnectionCountingPeer(("127.0.0.1", 0)) + thread: Final = threading.Thread(target=server.serve_forever, daemon=True) + thread.start() + yield f"http://127.0.0.1:{server.server_address[1]}" + server.shutdown() + server.server_close() + thread.join(timeout=10) + + +def _rebuilt_session_exchange(base_url: str) -> list[dict[str, JsonValue]]: + completed: Final = subprocess.run( + [sys.executable, "-P", "-c", REBUILT_SESSION_EXCHANGE, base_url, str(IDLE_SECONDS)], + env={ + **os.environ, + "AIOHTTP_KEEPALIVE_TIMEOUT": str(CONFIGURED_KEEPALIVE_SECONDS), + "AIOHTTP_SO_KEEPALIVE": "true", + }, + capture_output=True, + text=True, + timeout=60, + check=False, + ) + assert completed.returncode == 0, completed.stderr + return RESPONSES.validate_json(completed.stdout) + + +@pytest.mark.covers("sdk.aiohttp_transport.rebuilt_shared_session_keeps_configured_keepalive_timeout") +def test_rebuilt_shared_session_drops_idle_connection_after_configured_keepalive_timeout( + connection_counting_peer: str, +) -> None: + observed: Final = _rebuilt_session_exchange(connection_counting_peer) + assert observed == [{"connection": 1}, {"connection": 2}], observed diff --git a/tests/integration/streaming/test_file_content_streaming.py b/tests/integration/streaming/test_file_content_streaming.py new file mode 100644 index 00000000000..9ec7665307b --- /dev/null +++ b/tests/integration/streaming/test_file_content_streaming.py @@ -0,0 +1,48 @@ +import threading +import uuid +from collections.abc import Callable +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.wire import Reply, Request, wire_server + +STREAM_CHUNK_BYTES: Final = 1024 * 1024 +HEAD: Final = b"h" * STREAM_CHUNK_BYTES +TAIL: Final = b'{"custom_id": "tail", "response": {"status_code": 200}}\n' + + +def _file_content_gated_after_head(gate: threading.Event) -> Callable[[Request], Reply]: + def respond(_request: Request) -> Reply: + return Reply(content_type="application/octet-stream", chunks=(HEAD, TAIL), gate_after_first=gate) + + return respond + + +@pytest.mark.covers("streaming.file_content.body_reaches_client_before_upstream_finishes_sending") +def test_file_content_streams_the_first_megabyte_to_the_client_before_the_upstream_sends_the_rest( + gateway: Gateway, +) -> None: + file_id: Final = "file-" + uuid.uuid4().hex + gate: Final = threading.Event() + with gateway.scenario() as scenario, wire_server(_file_content_gated_after_head(gate)) as wire: + model: Final = scenario.model(api_base=wire.url + "/v1") + with gateway.client.stream( + "GET", + f"/v1/files/{file_id}/content", + params={"model": model}, + headers={"Authorization": f"Bearer {gateway.key}"}, + ) as response: + assert response.status_code == 200, response.read().decode() + chunks: Final = response.iter_bytes(chunk_size=STREAM_CHUNK_BYTES) + head: Final = next(chunks) + assert head == HEAD, f"First {len(head)} bytes differ from the upstream head before the gate was released" + gate.set() + rest: Final = b"".join(chunks) + assert rest == TAIL, rest + requests: Final = wire.drain() + assert len(requests) == 1, requests + assert requests[0].method == "GET", requests[0] + assert requests[0].target == f"/v1/files/{file_id}/content", requests[0].target + assert requests[0].headers["authorization"] == "Bearer integration-provider-key", requests[0].headers + assert requests[0].body == b"", requests[0].body diff --git a/tests/integration/streaming/test_stream_contracts.py b/tests/integration/streaming/test_stream_contracts.py index 7466a9454b8..bd89b869ef2 100644 --- a/tests/integration/streaming/test_stream_contracts.py +++ b/tests/integration/streaming/test_stream_contracts.py @@ -261,6 +261,87 @@ def test_messages_stream_completes_through_trailing_empty_choices_usage_chunk(ga ) +def reasoning_first_stream(identity: str) -> tuple[bytes, ...]: + usage: Final = { + "id": identity, + "object": "chat.completion.chunk", + "created": 1, + "model": "gpt-4o-mini", + "choices": [], + "usage": {"prompt_tokens": 11, "completion_tokens": 6, "total_tokens": 17}, + } + return ( + frame(identity, {"role": "assistant", "content": None, "reasoning_content": "Let me "}), + frame(identity, {"content": None, "reasoning_content": "think."}), + frame(identity, {"content": "Hello "}), + frame(identity, {"content": "there"}), + frame(identity, {}, finish="stop"), + b"data: " + json.dumps(usage).encode() + b"\n\n", + b"data: [DONE]\n\n", + ) + + +@pytest.mark.covers("streaming.messages_bridge.reasoning_content_only_chunks_open_a_thinking_block_first") +def test_messages_stream_opens_thinking_block_at_index_zero_for_reasoning_content_only_chunks( + gateway: Gateway, +) -> None: + identity: Final = "messages-reasoning-first-" + uuid.uuid4().hex + with ( + wire_server( + lambda request: Reply(content_type="text/event-stream", chunks=reasoning_first_stream(identity)) + ) as wire, + gateway.scenario() as scenario, + ): + model: Final = scenario.model(model="hosted_vllm/reasoning-model", api_base=wire.url + "/v1") + with gateway.client.stream( + "POST", + "/v1/messages", + json={ + "model": model, + "max_tokens": 64, + "stream": True, + "messages": [{"role": "user", "content": identity}], + }, + headers={"Authorization": f"Bearer {gateway.key}"}, + ) as response: + text: Final = response.read().decode() + assert response.status_code == 200, text + assert response.headers["content-type"].startswith("text/event-stream"), text + events: Final = tuple(json.loads(line) for line in sse_data_lines(text)) + blocks: Final = tuple( + (event["index"], event.get("content_block") or event["delta"]) + for event in events + if event["type"] in ("content_block_start", "content_block_delta") + ) + assert blocks == ( + (0, {"type": "thinking", "thinking": "", "signature": ""}), + (0, {"type": "thinking_delta", "thinking": "Let me "}), + (0, {"type": "thinking_delta", "thinking": "think."}), + (1, {"type": "text", "text": ""}), + (1, {"type": "text_delta", "text": "Hello "}), + (1, {"type": "text_delta", "text": "there"}), + ), text + assert tuple(event["type"] for event in events) == ( + "message_start", + "content_block_start", + "content_block_delta", + "content_block_delta", + "content_block_stop", + "content_block_start", + "content_block_delta", + "content_block_delta", + "content_block_stop", + "message_delta", + "message_stop", + ), text + message_delta: Final = next(event for event in events if event["type"] == "message_delta") + assert message_delta["usage"] == {"input_tokens": 11, "output_tokens": 6}, text + requests: Final = wire.drain() + assert len(requests) == 1 + outbound: Final = json.loads(requests[0].body) + assert outbound["stream"] is True and outbound["messages"] == [{"role": "user", "content": identity}], outbound + + @pytest.mark.covers("other.streaming.responses_bridge.empty_choices_chunks_complete_stream") def test_responses_stream_completes_through_empty_choices_metadata_and_usage_chunks(gateway: Gateway) -> None: identity: Final = "responses-empty-choices-" + uuid.uuid4().hex @@ -445,9 +526,7 @@ def test_primary_stream_with_empty_first_chunk_then_disconnect_falls_back_and_bi abort_after=2, ) ) as primary, - wire_server( - lambda request: Reply(content_type="text/event-stream", chunks=text_stream(identity)) - ) as fallback, + wire_server(lambda request: Reply(content_type="text/event-stream", chunks=text_stream(identity))) as fallback, ): config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) config["model_list"] = [ From 514bc181d65898b23692b8663084f98e1d123970 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 09:51:56 -0700 Subject: [PATCH 09/10] test(integration): regression tests for July cost tracking, budgeting and spend bugs (#42694) * test(integration): streamed Bedrock Messages usage cost equals the recorded spend (Pylon #6667) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): echoed cost-map model info is not persisted as deployment overrides (Pylon #6844) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): reset sweep runs on one pod per tick while replicas share the lease (Pylon #6521) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): bedrock post-call guardrail scans streamed Anthropic Messages tool use without 500 (Pylon #6503) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): realtime cached audio tokens bill at the audio cache-read rate (Pylon #6704) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): legacy GET /spend/logs returns at most the 10000 most recent rows and flags truncation (Pylon #6752) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): itemize Responses API cache write tokens as cache creation cost (Pylon #6454) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): migration entrypoint deploys pending migrations before proxy startup (Pylon #6649) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): opted-in team keys stop at the owner's personal budget (Pylon #6641) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): guardrail information stays in the spend log when the caller sends metadata on /v1/messages (Pylon #6614) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): key model allowlist is enforced on Bedrock passthrough routes (Pylon #6419) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): JWT mapped key backfills a null user email from token claims (Pylon #6266) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): scheduled budget reset recovers from a transient DB transport failure (Pylon #6582) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): team key lists models granted through a team access group (Pylon #6044) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): JWT subject without team claim lands in the configured default team (Pylon #5895) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): end-user spend lands for a key without user_id when the auth cache is Redis (Pylon #6021) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): stale-low redis counter still blocks team member over budget (Pylon #5824) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): prompt-carrying spend rows are written in byte-bounded statements (Pylon #6083) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): logs UI session_total_spend sums every round of a multi-round session (Pylon #5928) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): config.yaml guardrails are served by the guardrail usage detail and overview (Pylon #5813) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): plain chat request skips the object permission lookup (Pylon #5965) Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): register july accounting regression contracts Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): isolate cost map override clear on owned proxy Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): make reset lease claim and db relay refusal deterministic Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): poll pg_stat settle, bound unbanned relay refusals, clear reset lease on teardown Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): bound relay refusals so the budget sweep can reconnect Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- tests/integration/_support/database.py | 6 +- tests/integration/_support/database_relay.py | 95 +++++++++ .../test_access_group_model_listing.py | 51 +++++ .../test_bedrock_passthrough_model_access.py | 65 ++++++ .../test_jwt_default_team_provisioning.py | 86 ++++++++ .../test_jwt_mapped_key_email_backfill.py | 112 ++++++++++ .../test_object_permission_lookup.py | 74 +++++++ .../database/test_migration_entrypoint.py | 99 +++++++++ .../test_guardrail_usage_config_guardrail.py | 77 +++++++ .../observability/test_guardrail_effects.py | 192 +++++++++++++++++- .../pricing/test_configured_prices.py | 45 ++++ .../test_realtime_cached_audio_pricing.py | 131 ++++++++++++ .../integration/spend/test_cache_and_quota.py | 106 +++++++++- .../test_end_user_spend_without_proxy_user.py | 29 +++ .../spend/test_legacy_spend_logs_row_cap.py | 46 +++++ .../spend/test_messages_stream_usage_cost.py | 176 ++++++++++++++++ .../test_reset_budget_leader_election.py | 75 +++++++ .../test_responses_cache_write_itemization.py | 98 +++++++++ .../spend/test_session_total_spend.py | 98 +++++++++ .../spend/test_spend_log_write_batching.py | 75 +++++++ .../spend/test_team_member_spend.py | 45 ++++ .../spend/test_user_budget_on_team_keys.py | 54 +++++ 22 files changed, 1830 insertions(+), 5 deletions(-) create mode 100644 tests/integration/_support/database_relay.py create mode 100644 tests/integration/authorization/test_access_group_model_listing.py create mode 100644 tests/integration/authorization/test_bedrock_passthrough_model_access.py create mode 100644 tests/integration/authorization/test_jwt_default_team_provisioning.py create mode 100644 tests/integration/authorization/test_jwt_mapped_key_email_backfill.py create mode 100644 tests/integration/authorization/test_object_permission_lookup.py create mode 100644 tests/integration/database/test_migration_entrypoint.py create mode 100644 tests/integration/management/test_guardrail_usage_config_guardrail.py create mode 100644 tests/integration/pricing/test_realtime_cached_audio_pricing.py create mode 100644 tests/integration/spend/test_end_user_spend_without_proxy_user.py create mode 100644 tests/integration/spend/test_legacy_spend_logs_row_cap.py create mode 100644 tests/integration/spend/test_messages_stream_usage_cost.py create mode 100644 tests/integration/spend/test_reset_budget_leader_election.py create mode 100644 tests/integration/spend/test_responses_cache_write_itemization.py create mode 100644 tests/integration/spend/test_session_total_spend.py create mode 100644 tests/integration/spend/test_spend_log_write_batching.py create mode 100644 tests/integration/spend/test_user_budget_on_team_keys.py diff --git a/tests/integration/_support/database.py b/tests/integration/_support/database.py index 283d26a632e..e7f0ebdf603 100644 --- a/tests/integration/_support/database.py +++ b/tests/integration/_support/database.py @@ -8,7 +8,9 @@ from pydantic import JsonValue, TypeAdapter ROWS: Final = TypeAdapter(list[dict[str, JsonValue]]) -def read_rows(query: str, parameters: tuple[str, ...]) -> list[dict[str, JsonValue]]: - with psycopg.connect(os.environ["DATABASE_URL"], row_factory=dict_row) as connection: +def read_rows( + query: str, parameters: tuple[str, ...], *, database_url: str | None = None +) -> list[dict[str, JsonValue]]: + with psycopg.connect(database_url or os.environ["DATABASE_URL"], row_factory=dict_row) as connection: connection.execute("SET TRANSACTION READ ONLY") return ROWS.validate_python(connection.execute(query, parameters).fetchall()) diff --git a/tests/integration/_support/database_relay.py b/tests/integration/_support/database_relay.py new file mode 100644 index 00000000000..46f3e17af13 --- /dev/null +++ b/tests/integration/_support/database_relay.py @@ -0,0 +1,95 @@ +import asyncio +import socket +import threading +from collections.abc import Generator +from contextlib import contextmanager +from typing import Final +from urllib.parse import urlsplit, urlunsplit + +from pydantic import TypeAdapter + +PORT: Final = TypeAdapter(int) + + +def _free_port() -> int: + with socket.socket() as reserve: + reserve.bind(("127.0.0.1", 0)) + return PORT.validate_python(reserve.getsockname()[1]) + + +class DatabaseRelay: + def __init__(self, upstream_host: str, upstream_port: int, trigger: bytes) -> None: + self.port: Final = _free_port() + self._upstream_host: Final = upstream_host + self._upstream_port: Final = upstream_port + self._trigger: Final = trigger + self._loop: Final = asyncio.new_event_loop() + self._armed: Final = threading.Event() + self.tripped: Final = threading.Event() + self.refused = 0 + self._writers: tuple[asyncio.StreamWriter, ...] = () + self._ready: Final = threading.Event() + self._thread: Final = threading.Thread(target=self._run, daemon=True) + + def arm(self) -> None: + self._armed.set() + + def start(self) -> None: + self._thread.start() + assert self._ready.wait(10), "Database relay did not start" + + def stop(self) -> None: + self._loop.call_soon_threadsafe(self._loop.stop) + self._thread.join(10) + + def _run(self) -> None: + asyncio.set_event_loop(self._loop) + self._loop.run_until_complete(asyncio.start_server(self._serve, "127.0.0.1", self.port)) + self._ready.set() + self._loop.run_forever() + + def _drop_all(self) -> None: + for writer in self._writers: + writer.close() + self._writers = () + + async def _serve(self, client_reader: asyncio.StreamReader, client_writer: asyncio.StreamWriter) -> None: + if self.tripped.is_set() and self.refused < 5: + self.refused += 1 + client_writer.close() + return + server_reader, server_writer = await asyncio.open_connection(self._upstream_host, self._upstream_port) + self._writers = (*self._writers, client_writer, server_writer) + + async def forward(reader: asyncio.StreamReader, writer: asyncio.StreamWriter, inspect: bool) -> None: + try: + while chunk := await reader.read(65536): + if inspect and self._armed.is_set() and not self.tripped.is_set() and self._trigger in chunk: + self.tripped.set() + self._drop_all() + return + writer.write(chunk) + await writer.drain() + except (ConnectionError, asyncio.IncompleteReadError): + return + finally: + writer.close() + + await asyncio.gather( + forward(client_reader, server_writer, True), + forward(server_reader, client_writer, False), + ) + + +@contextmanager +def database_relay(database_url: str, trigger: bytes) -> Generator[tuple[DatabaseRelay, str]]: + parts: Final = urlsplit(database_url) + assert parts.hostname is not None and parts.port is not None, database_url + relay: Final = DatabaseRelay(parts.hostname, parts.port, trigger) + relay.start() + credentials: Final = f"{parts.username}:{parts.password}@" if parts.username else "" + relayed: Final = urlunsplit(parts._replace(netloc=f"{credentials}127.0.0.1:{relay.port}")) + try: + yield relay, relayed + finally: + relay.stop() diff --git a/tests/integration/authorization/test_access_group_model_listing.py b/tests/integration/authorization/test_access_group_model_listing.py new file mode 100644 index 00000000000..9b37cc8f232 --- /dev/null +++ b/tests/integration/authorization/test_access_group_model_listing.py @@ -0,0 +1,51 @@ +import uuid +from collections.abc import Iterator +from contextlib import contextmanager +from typing import Final + +import httpx +import pytest + +from tests.integration._support.client import Gateway, eventually, object_value, string_value + + +@contextmanager +def _team_access_group(gateway: Gateway, team_id: str, model: str) -> Iterator[str]: + created: Final = gateway.request( + "POST", + "/v1/access_group", + { + "access_group_name": f"integration-{uuid.uuid4().hex}", + "access_model_names": [model], + "assigned_team_ids": [team_id], + }, + ) + assert created.status_code == 201, created.text + identity: Final = string_value(created.json()["access_group_id"]) + try: + yield identity + finally: + deleted: Final = gateway.request("DELETE", f"/v1/access_group/{identity}") + assert deleted.status_code == 204, deleted.text + + +def _listed_model_ids(response: httpx.Response) -> tuple[str, ...]: + entries: Final = response.json()["data"] + assert isinstance(entries, list), response.text + return tuple(string_value(object_value(entry)["id"]) for entry in entries) + + +@pytest.mark.covers("authorization.access_groups.team_key_lists_models_granted_through_team_access_group") +def test_team_key_lists_models_granted_through_team_access_group(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + team_id: Final = scenario.team(models=["no-default-models"]) + with _team_access_group(gateway, team_id, model): + key: Final = scenario.key(team_id=team_id) + response: Final = eventually( + lambda: gateway.request("GET", "/v1/models", key=key), + lambda value: value.status_code == 200 and _listed_model_ids(value) == (model,), + return_last_on_timeout=True, + ) + assert response.status_code == 200, response.text + assert _listed_model_ids(response) == (model,), response.text diff --git a/tests/integration/authorization/test_bedrock_passthrough_model_access.py b/tests/integration/authorization/test_bedrock_passthrough_model_access.py new file mode 100644 index 00000000000..c1b54fe7e38 --- /dev/null +++ b/tests/integration/authorization/test_bedrock_passthrough_model_access.py @@ -0,0 +1,65 @@ +import json +from typing import Final + +import pytest +from integration._support.client import Gateway +from integration._support.upstream import _aws_event_frame +from integration._support.wire import Reply, Request, wire_server + +_MODEL_ID: Final = "anthropic.claude-sonnet-5-v1:0" +_ACTIONS: Final = ("converse", "invoke", "converse-stream", "invoke-with-response-stream") +_REQUEST_BODY: Final = {"messages": [{"role": "user", "content": [{"text": "synthetic passthrough allowlist"}]}]} +_CONVERSE_RESPONSE: Final = json.dumps( + { + "output": {"message": {"role": "assistant", "content": [{"text": "bedrock allowlist control"}]}}, + "stopReason": "end_turn", + "usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15}, + "metrics": {"latencyMs": 1}, + } +).encode() +_STREAM_BYTES: Final = b"".join( + _aws_event_frame(kind, payload, "sc", "u") + for kind, payload in ( + ("messageStart", {"role": "assistant"}), + ("contentBlockDelta", {"delta": {"text": "bedrock allowlist control"}, "contentBlockIndex": 0}), + ("messageStop", {"stopReason": "end_turn"}), + ("metadata", {"usage": {"inputTokens": 11, "outputTokens": 4, "totalTokens": 15}}), + ) +) + + +def bedrock_peer(request: Request) -> Reply: + assert request.method == "POST", request.target + assert json.loads(request.body)["messages"] == _REQUEST_BODY["messages"], request.body + if request.target.endswith("-stream"): + return Reply(body=_STREAM_BYTES, content_type="application/vnd.amazon.eventstream") + return Reply(body=_CONVERSE_RESPONSE) + + +@pytest.mark.covers("authz.key_models.bedrock_passthrough_route_model_is_enforced") +def test_key_scoped_to_one_model_cannot_call_another_through_bedrock_passthrough_routes(gateway: Gateway) -> None: + with wire_server(bedrock_peer) as wire, gateway.scenario() as scenario: + allowed: Final = scenario.model( + model=f"bedrock/{_MODEL_ID}", + api_base=wire.url, + aws_access_key_id="AKIASCRIPTEDPROVIDER", + aws_secret_access_key="scripted-secret", + aws_region_name="us-east-1", + ) + denied: Final = scenario.model( + model=f"bedrock/{_MODEL_ID}", + api_base=wire.url, + aws_access_key_id="AKIASCRIPTEDPROVIDER", + aws_secret_access_key="scripted-secret", + aws_region_name="us-east-1", + ) + key: Final = scenario.key(models=[allowed]) + for action in _ACTIONS: + response: Final = gateway.request("POST", f"/bedrock/model/{denied}/{action}", _REQUEST_BODY, key=key) + assert response.status_code == 403, f"{action}: {response.status_code} {response.text}" + assert response.json()["error"]["type"] == "key_model_access_denied", f"{action}: {response.text}" + assert wire.drain() == (), f"{action} reached the provider: {response.text}" + for action in _ACTIONS: + served: Final = gateway.request("POST", f"/bedrock/model/{allowed}/{action}", _REQUEST_BODY, key=key) + assert served.status_code == 200, f"{action}: {served.status_code} {served.text}" + assert tuple(request.target for request in wire.drain()) == (f"/model/{_MODEL_ID}/{action}",), served.text diff --git a/tests/integration/authorization/test_jwt_default_team_provisioning.py b/tests/integration/authorization/test_jwt_default_team_provisioning.py new file mode 100644 index 00000000000..04715d128b0 --- /dev/null +++ b/tests/integration/authorization/test_jwt_default_team_provisioning.py @@ -0,0 +1,86 @@ +import json +import time +import uuid +from pathlib import Path +from typing import Final + +import jwt +import pytest +import yaml +from cryptography.hazmat.primitives.asymmetric import rsa + +from tests.integration._support.client import Gateway, eventually +from tests.integration._support.database import read_rows +from tests.integration._support.process import owned_proxy +from tests.integration._support.wire import Reply, Request, wire_server + +KEY_ID: Final = "integration-jwt-signing-key" +TEAM_BUDGET: Final = 25.0 + + +def _jwks_reply(public_jwk: str) -> Reply: + return Reply(body=json.dumps({"keys": [{**json.loads(public_jwk), "kid": KEY_ID}]}).encode()) + + +@pytest.mark.covers("authorization.jwt.new_subject_without_team_claim_joins_default_team") +def test_jwt_subject_without_team_claim_is_provisioned_into_configured_default_team( + gateway: Gateway, tmp_path: Path +) -> None: + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + public_jwk: Final = jwt.algorithms.RSAAlgorithm.to_jwk(private_key.public_key()) + + def respond(request: Request) -> Reply: + assert request.method == "GET", request + return _jwks_reply(public_jwk) + + with wire_server(respond) as jwks, gateway.scenario() as scenario: + team: Final = scenario.team() + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["general_settings"] = { + **config["general_settings"], + "enable_jwt_auth": True, + "litellm_jwtauth": {"user_id_jwt_field": "sub", "user_id_upsert": True}, + } + config["litellm_settings"] = { + **config["litellm_settings"], + "default_internal_user_params": { + "user_role": "internal_user", + "teams": [{"team_id": team, "user_role": "user", "max_budget_in_team": TEAM_BUDGET}], + }, + } + path: Final = tmp_path / "jwt_default_team.yaml" + path.write_text(yaml.safe_dump(config)) + subject: Final = f"integration-jwt-{uuid.uuid4().hex}" + token: Final = jwt.encode( + {"sub": subject, "iat": int(time.time()), "exp": int(time.time()) + 300}, + private_key, + algorithm="RS256", + headers={"kid": KEY_ID}, + ) + with owned_proxy(gateway, tmp_path, {"JWT_PUBLIC_KEY_URL": jwks.url}, config=path) as candidate: + model: Final = scenario.model() + scenario.cleanups.callback(scenario.delete_user, subject) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "default team control"}]}, + key=token, + ) + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"] == ( + "Hello! This is a mock response from the fake OpenAI endpoint." + ), response.text + assert read_rows( + 'SELECT user_id, user_role, teams FROM "LiteLLM_UserTable" WHERE user_id = %s', (subject,) + ) == [{"user_id": subject, "user_role": "internal_user", "teams": [team]}] + memberships: Final = eventually( + lambda: read_rows( + 'SELECT m.team_id, b.max_budget FROM "LiteLLM_TeamMembership" m ' + 'JOIN "LiteLLM_BudgetTable" b ON b.budget_id = m.budget_id WHERE m.user_id = %s', + (subject,), + ), + lambda rows: len(rows) == 1, + ) + assert memberships == [{"team_id": team, "max_budget": TEAM_BUDGET}] + roster: Final = read_rows('SELECT members_with_roles FROM "LiteLLM_TeamTable" WHERE team_id = %s', (team,)) + assert {"user_id": subject, "role": "user", "user_email": None} in roster[0]["members_with_roles"], roster diff --git a/tests/integration/authorization/test_jwt_mapped_key_email_backfill.py b/tests/integration/authorization/test_jwt_mapped_key_email_backfill.py new file mode 100644 index 00000000000..47c8ad5f115 --- /dev/null +++ b/tests/integration/authorization/test_jwt_mapped_key_email_backfill.py @@ -0,0 +1,112 @@ +import json +import time +import uuid +from hashlib import sha256 +from pathlib import Path +from typing import Final + +import jwt +import pytest +from cryptography.hazmat.primitives.asymmetric import rsa +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.process import owned_proxy +from integration._support.wire import Reply, Request, wire_server +from jwt.algorithms import RSAAlgorithm + +AUDIENCE: Final = "litellm-integration" +KEY_ID: Final = "integration-signing-key" +CLIENT_CLAIM: Final = "client_id" + + +def _proxy_config(directory: Path, model: str, upstream_url: str) -> Path: + config: Final = directory / "jwt_mapped_key_config.yaml" + config.write_text( + json.dumps( + { + "model_list": [ + { + "model_name": model, + "litellm_params": { + "model": "openai/" + model, + "api_base": upstream_url + "/v1", + "api_key": "sk-upstream", + }, + } + ], + "general_settings": { + "master_key": "os.environ/LITELLM_MASTER_KEY", + "database_url": "os.environ/DATABASE_URL", + "store_model_in_db": True, + "proxy_batch_write_at": 1, + "proxy_batch_polling_interval": 1, + "enable_jwt_auth": True, + "litellm_jwtauth": { + "user_id_jwt_field": "sub", + "user_email_jwt_field": "email", + "virtual_key_claim_field": CLIENT_CLAIM, + }, + }, + "router_settings": {"disable_cooldowns": True}, + } + ) + ) + return config + + +def _signed_token(private_key: rsa.RSAPrivateKey, user_id: str, email: str, client_id: str) -> str: + now: Final = int(time.time()) + return jwt.encode( + {"sub": user_id, "email": email, CLIENT_CLAIM: client_id, "aud": AUDIENCE, "iat": now, "exp": now + 300}, + private_key, + algorithm="RS256", + headers={"kid": KEY_ID}, + ) + + +@pytest.mark.covers("authorization.jwt.mapped_key_backfills_null_user_email_from_claims") +def test_jwt_mapped_key_request_backfills_null_user_email_from_token_claims(gateway: Gateway, tmp_path: Path) -> None: + private_key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + public_jwk: Final = json.loads(RSAAlgorithm.to_jwk(private_key.public_key())) + jwks: Final = json.dumps({"keys": [{**public_jwk, "kid": KEY_ID, "use": "sig", "alg": "RS256"}]}).encode() + + def respond(request: Request) -> Reply: + assert request.target == "/jwks", request + return Reply(body=jwks) + + model: Final = "integration-jwt-" + uuid.uuid4().hex + with wire_server(respond) as issuer: + config: Final = _proxy_config(tmp_path, model, gateway.upstream_url) + overrides: Final = {"JWT_PUBLIC_KEY_URL": issuer.url + "/jwks", "JWT_AUDIENCE": AUDIENCE} + with owned_proxy(gateway, tmp_path, overrides, config=config) as candidate, candidate.scenario() as scenario: + user: Final = scenario.user() + key: Final = scenario.key(user_id=user, models=[model]) + client_id: Final = "integration-client-" + uuid.uuid4().hex + mapping: Final = candidate.post( + "/jwt/key/mapping/new", {"jwt_claim_name": CLIENT_CLAIM, "jwt_claim_value": client_id, "key": key} + ) + scenario.cleanups.callback(candidate.post, "/jwt/key/mapping/delete", {"id": mapping["id"]}) + assert read_rows('SELECT user_email FROM "LiteLLM_UserTable" WHERE user_id = %s', (user,)) == [ + {"user_email": None} + ] + email: Final = f"{user}@integration.example" + token: Final = _signed_token(private_key, user, email, client_id) + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "jwt email backfill control"}]}, + key=token, + ) + assert response.status_code == 200, response.text + assert read_rows('SELECT user_email FROM "LiteLLM_UserTable" WHERE user_id = %s', (user,)) == [ + {"user_email": email} + ] + spend_rows: Final = eventually( + lambda: read_rows( + 'SELECT api_key, "user" FROM "LiteLLM_SpendLogs" WHERE request_id = %s', + (str(response.json()["id"]),), + ), + lambda rows: len(rows) == 1, + seconds=70, + ) + assert spend_rows == [{"api_key": sha256(key.encode()).hexdigest(), "user": user}] diff --git a/tests/integration/authorization/test_object_permission_lookup.py b/tests/integration/authorization/test_object_permission_lookup.py new file mode 100644 index 00000000000..21fba30c165 --- /dev/null +++ b/tests/integration/authorization/test_object_permission_lookup.py @@ -0,0 +1,74 @@ +import time +from typing import Final + +import pytest + +from tests.integration._support.client import Gateway, eventually +from tests.integration._support.database import read_rows + +PLAIN_REQUESTS: Final = 10 +STATS_FLUSH_WINDOW_SECONDS: Final = 11.0 + + +def _object_permission_reads() -> int: + rows: Final = read_rows( + "SELECT seq_scan + idx_scan AS reads FROM pg_stat_user_tables WHERE relname = %s", + ("LiteLLM_ObjectPermissionTable",), + ) + reads: Final = rows[0]["reads"] + assert isinstance(reads, int), rows + return reads + + +def _settled_object_permission_reads(previous: int, unchanged_since: float) -> int: + changed: Final = eventually( + lambda: _object_permission_reads() != previous, + lambda drifted: drifted, + seconds=STATS_FLUSH_WINDOW_SECONDS - (time.monotonic() - unchanged_since), + return_last_on_timeout=True, + ) + if not changed: + return previous + return _settled_object_permission_reads(_object_permission_reads(), time.monotonic()) + + +def _assert_plain_chat_served(gateway: Gateway, model: str, key: str) -> None: + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "no vector stores"}]}, + key=key, + ) + assert response.status_code == 200, response.text + assert response.json()["choices"][0]["message"]["content"] == ( + "Hello! This is a mock response from the fake OpenAI endpoint." + ) + + +def _assert_forbidden_vector_store_denied(gateway: Gateway, model: str, key: str) -> None: + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "forbidden"}], "vector_store_ids": ["vs_forbidden"]}, + key=key, + ) + assert response.status_code == 401, response.text + assert response.json()["error"]["type"] == "key_vector_store_access_denied", response.text + + +@pytest.mark.covers("authorization.vector_store.plain_request_skips_object_permission_lookup") +def test_chat_request_without_vector_stores_does_not_read_object_permission_table(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model() + key: Final = scenario.key(models=[model], object_permission={"vector_stores": ["vs_allowed"]}) + _assert_plain_chat_served(gateway, model, key) + before_control: Final = _object_permission_reads() + _assert_forbidden_vector_store_denied(gateway, model, key) + eventually(_object_permission_reads, lambda reads: reads > before_control, seconds=15) + baseline: Final = _settled_object_permission_reads(_object_permission_reads(), time.monotonic()) + for _ in range(PLAIN_REQUESTS): + _assert_plain_chat_served(gateway, model, key) + after: Final = _settled_object_permission_reads(_object_permission_reads(), time.monotonic()) + assert after - baseline < PLAIN_REQUESTS, ( + f"{PLAIN_REQUESTS} plain chat requests added {after - baseline} object permission reads" + ) diff --git a/tests/integration/database/test_migration_entrypoint.py b/tests/integration/database/test_migration_entrypoint.py new file mode 100644 index 00000000000..19c480b9531 --- /dev/null +++ b/tests/integration/database/test_migration_entrypoint.py @@ -0,0 +1,99 @@ +import os +import shutil +import subprocess +import sys +import uuid +from collections.abc import Iterator +from contextlib import contextmanager +from pathlib import Path +from typing import Final +from urllib.parse import urlsplit, urlunsplit + +import psycopg +import pytest +from integration._support.client import Gateway +from integration._support.process import owned_proxy +from psycopg import sql +from psycopg.rows import dict_row + +REPO_ROOT: Final = Path(__file__).resolve().parents[3] +PRISMA_DIR: Final = REPO_ROOT / "litellm-proxy-extras" / "litellm_proxy_extras" +MISSING_MIGRATION: Final = "20260626120000_add_mcp_tool_search_enabled" +SHIPPED_MIGRATIONS: Final = tuple(sorted(path.name for path in (PRISMA_DIR / "migrations").iterdir() if path.is_dir())) + + +@contextmanager +def fresh_database() -> Iterator[str]: + name: Final = f"integration_upgrade_{uuid.uuid4().hex}" + admin_url: Final = os.environ["DATABASE_URL"] + parsed: Final = urlsplit(admin_url) + with psycopg.connect(admin_url, autocommit=True) as admin: + admin.execute(sql.SQL("CREATE DATABASE {}").format(sql.Identifier(name))) + try: + yield urlunsplit(parsed._replace(path=f"/{name}")) + finally: + admin.execute(sql.SQL("DROP DATABASE {} WITH (FORCE)").format(sql.Identifier(name))) + + +def deploy_older_schema(database_url: str, directory: Path) -> None: + older: Final = directory / "older-release" + (older / "migrations").mkdir(parents=True) + shutil.copy(PRISMA_DIR / "schema.prisma", older / "schema.prisma") + shutil.copy(PRISMA_DIR / "migrations" / "migration_lock.toml", older / "migrations" / "migration_lock.toml") + for name in (name for name in SHIPPED_MIGRATIONS if name < MISSING_MIGRATION): + shutil.copytree(PRISMA_DIR / "migrations" / name, older / "migrations" / name) + subprocess.run( + [sys.executable, "-I", "-m", "prisma", "migrate", "deploy", "--schema", str(older / "schema.prisma")], + check=True, + capture_output=True, + text=True, + timeout=300, + env={**os.environ, "DATABASE_URL": database_url}, + ) + + +def applied_migrations(database_url: str) -> tuple[str, ...]: + with psycopg.connect(database_url, row_factory=dict_row) as connection: + rows: Final = connection.execute( + 'SELECT migration_name FROM "_prisma_migrations" ' + "WHERE finished_at IS NOT NULL AND rolled_back_at IS NULL ORDER BY migration_name" + ).fetchall() + return tuple(str(row["migration_name"]) for row in rows) + + +def object_permission_columns(database_url: str) -> tuple[str, ...]: + with psycopg.connect(database_url, row_factory=dict_row) as connection: + rows: Final = connection.execute( + "SELECT column_name FROM information_schema.columns " + "WHERE table_name = 'LiteLLM_ObjectPermissionTable' AND column_name = 'mcp_tool_search_enabled'" + ).fetchall() + return tuple(str(row["column_name"]) for row in rows) + + +@pytest.mark.covers("other.database.migrations.entrypoint_deploys_pending_migrations_before_startup") +def test_migration_entrypoint_upgrades_an_older_schema_so_the_proxy_serves_mcp_tools( + gateway: Gateway, tmp_path: Path +) -> None: + with fresh_database() as database_url: + deploy_older_schema(database_url, tmp_path) + assert object_permission_columns(database_url) == () + assert applied_migrations(database_url) == tuple( + name for name in SHIPPED_MIGRATIONS if name < MISSING_MIGRATION + ) + entrypoint: Final = subprocess.run( + [sys.executable, "-I", "-m", "litellm.proxy.prisma_migration"], + capture_output=True, + text=True, + timeout=300, + cwd=REPO_ROOT, + env={**os.environ, "DATABASE_URL": database_url}, + ) + assert entrypoint.returncode == 0, entrypoint.stdout + entrypoint.stderr + assert object_permission_columns(database_url) == ("mcp_tool_search_enabled",), entrypoint.stdout + assert applied_migrations(database_url) == SHIPPED_MIGRATIONS, entrypoint.stdout + with owned_proxy( + gateway, tmp_path, {"DATABASE_URL": database_url, "DISABLE_SCHEMA_UPDATE": "true"} + ) as upgraded: + tools: Final = upgraded.request("GET", "/mcp-rest/tools/list") + assert tools.status_code == 200, tools.text + assert tools.json()["tools"] == [], tools.text diff --git a/tests/integration/management/test_guardrail_usage_config_guardrail.py b/tests/integration/management/test_guardrail_usage_config_guardrail.py new file mode 100644 index 00000000000..b9ba795b89b --- /dev/null +++ b/tests/integration/management/test_guardrail_usage_config_guardrail.py @@ -0,0 +1,77 @@ +import uuid +from pathlib import Path +from typing import Final + +import pytest +import yaml +from pydantic import JsonValue + +from tests.integration._support.client import Gateway, object_value, string_value +from tests.integration._support.process import owned_proxy + +WINDOW: Final = {"start_date": "2026-01-01", "end_date": "2026-01-07"} + + +def listed_guardrail(gateway: Gateway, guardrail_name: str) -> dict[str, JsonValue]: + listed: Final = gateway.get("/v2/guardrails/list")["guardrails"] + assert isinstance(listed, list), listed + matches: Final = tuple(object_value(row) for row in listed if object_value(row)["guardrail_name"] == guardrail_name) + assert len(matches) == 1, f"{guardrail_name} appears {len(matches)} times in {listed}" + return matches[0] + + +@pytest.mark.covers("mgmt.guardrails.usage.config_yaml_guardrail_has_detail_and_overview_row") +def test_config_yaml_guardrail_is_served_by_usage_detail_and_overview(gateway: Gateway, tmp_path: Path) -> None: + guardrail_name: Final = "tool-permission-" + uuid.uuid4().hex + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": guardrail_name, + "litellm_params": { + "guardrail": "tool_permission", + "mode": "post_call", + "default_on": False, + "rules": [{"id": "deny_delete", "tool_name": "(?i)^.*(delete|drop).*", "decision": "deny"}], + "default_action": "allow", + "on_disallowed_action": "block", + }, + "guardrail_info": {"type": "Tool Permission", "description": "declared in config.yaml"}, + } + ] + path: Final = tmp_path / "guardrail.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate: + guardrail_id: Final = string_value(listed_guardrail(candidate, guardrail_name)["guardrail_id"]) + detail: Final = candidate.request("GET", f"/guardrails/usage/detail/{guardrail_id}", params=WINDOW) + assert detail.status_code == 200, detail.text + body: Final = object_value(detail.json()) + assert { + "guardrail_id": body["guardrail_id"], + "guardrail_name": body["guardrail_name"], + "provider": body["provider"], + "type": body["type"], + "description": body["description"], + "requestsEvaluated": body["requestsEvaluated"], + "failRate": body["failRate"], + } == { + "guardrail_id": guardrail_id, + "guardrail_name": guardrail_name, + "provider": "tool_permission", + "type": "Tool Permission", + "description": "declared in config.yaml", + "requestsEvaluated": 0, + "failRate": 0.0, + }, detail.text + overview: Final = candidate.request("GET", "/guardrails/usage/overview", params=WINDOW) + assert overview.status_code == 200, overview.text + rows: Final = object_value(overview.json())["rows"] + assert isinstance(rows, list), overview.text + config_rows: Final = tuple(object_value(row) for row in rows if object_value(row)["id"] == guardrail_id) + assert len(config_rows) == 1, overview.text + assert (config_rows[0]["name"], config_rows[0]["provider"], config_rows[0]["requestsEvaluated"]) == ( + guardrail_name, + "tool_permission", + 0, + ), overview.text + missing: Final = candidate.request("GET", f"/guardrails/usage/detail/{uuid.uuid4()}", params=WINDOW) + assert missing.status_code == 404, missing.text diff --git a/tests/integration/observability/test_guardrail_effects.py b/tests/integration/observability/test_guardrail_effects.py index f96cae593dc..4fac42a796d 100644 --- a/tests/integration/observability/test_guardrail_effects.py +++ b/tests/integration/observability/test_guardrail_effects.py @@ -5,7 +5,8 @@ from typing import Final import pytest import yaml -from integration._support.client import Gateway +from integration._support.client import Gateway, eventually, object_value +from integration._support.database import read_rows from integration._support.mcp import mcp_peer, register_mcp, tool_names from integration._support.process import owned_proxy from integration._support.wire import Reply, Request, wire_server @@ -88,6 +89,91 @@ def test_guardrail_rewrites_system_and_user_in_actual_anthropic_request(gateway: assert len(policy.drain()) == len(upstream.drain()) == 1 +@pytest.mark.covers("other.observability.guardrails.anthropic_messages_caller_metadata_keeps_guardrail_spend_log") +def test_anthropic_messages_with_caller_metadata_keeps_guardrail_information_in_spend_log( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + prompt: Final = "synthetic allowed prompt " + identity + caller_metadata: Final = {"user_id": "device-account-session"} + + def guardrail(request: Request) -> Reply: + assert request.target == "/beta/litellm_basic_guardrail_api" + assert json.loads(request.body)["texts"] == [prompt] + return Reply(body=json.dumps({"action": "NONE"}).encode()) + + def provider(request: Request) -> Reply: + assert request.target == "/v1/messages" + body: Final = json.loads(request.body) + assert body["messages"] == [{"role": "user", "content": prompt}] + assert body["metadata"] == caller_metadata + return Reply( + body=json.dumps( + { + "id": identity, + "type": "message", + "role": "assistant", + "model": "claude-sonnet-4-5-20250929", + "content": [{"type": "text", "text": "permitted response"}], + "stop_reason": "end_turn", + "stop_sequence": None, + "usage": {"input_tokens": 11, "output_tokens": 4}, + } + ).encode() + ) + + with wire_server(guardrail) as policy, wire_server(provider) as upstream: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "generic_guardrail_api", + "mode": "pre_call", + "default_on": True, + "api_base": policy.url, + "api_key": "synthetic-guardrail-key", + }, + } + ] + path: Final = tmp_path / "caller-metadata.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", api_base=upstream.url, api_key="synthetic-anthropic-key" + ) + response: Final = candidate.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 16, + "messages": [{"role": "user", "content": prompt}], + "metadata": caller_metadata, + }, + ) + assert response.status_code == 200, response.text + assert response.json()["content"] == [{"type": "text", "text": "permitted response"}], response.text + assert response.headers["x-litellm-applied-guardrails"] == identity, dict(response.headers) + assert len(policy.drain()) == len(upstream.drain()) == 1 + rows: Final = eventually( + lambda: read_rows( + 'SELECT call_type, metadata FROM "LiteLLM_SpendLogs" WHERE model_group=%s', + (model,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert rows[0]["call_type"] == "anthropic_messages", rows[0] + saved: Final = object_value(rows[0]["metadata"]) + entries: Final = saved["guardrail_information"] + assert isinstance(entries, list) and len(entries) == 1, saved + entry: Final = object_value(entries[0]) + assert entry["guardrail_name"] == identity, saved + assert entry["guardrail_mode"] == "pre_call", saved + assert entry["guardrail_status"] == "success", saved + + @pytest.mark.covers("other.observability.guardrails.denial_prevents_provider_with_allowed_control") def test_guardrail_denial_prevents_provider_and_preserves_allowed_control(gateway: Gateway, tmp_path: Path) -> None: identity: Final = "guardrail" + uuid.uuid4().hex @@ -242,6 +328,110 @@ def test_bedrock_passthrough_converse_guardrail_ignores_denied_term_in_tool_defi assert [json.loads(request.body)["texts"] for request in policy.drain()] == [[allowed], [denied]] +@pytest.mark.covers("other.observability.guardrails.bedrock_post_call_scans_streamed_anthropic_messages_tool_use") +def test_bedrock_guardrail_streams_anthropic_messages_tool_use_instead_of_chunk_builder_500( + gateway: Gateway, tmp_path: Path +) -> None: + identity: Final = "guardrail" + uuid.uuid4().hex + guardrail_id: Final = "synthetic" + uuid.uuid4().hex[:8] + spoken: Final = "Checking the forecast" + frames: Final = ( + 'event: message_start\ndata: {"type": "message_start", "message": {"id": "msg_synthetic", "type": "message", ' + '"role": "assistant", "model": "claude-sonnet-4-5-20250929", "content": [], "stop_reason": null, ' + '"stop_sequence": null, "usage": {"input_tokens": 11, "output_tokens": 1}}}\n\n', + 'event: content_block_start\ndata: {"type": "content_block_start", "index": 0, ' + '"content_block": {"type": "text", "text": ""}}\n\n', + 'event: content_block_delta\ndata: {"type": "content_block_delta", "index": 0, ' + f'"delta": {{"type": "text_delta", "text": "{spoken}"}}}}\n\n', + 'event: content_block_stop\ndata: {"type": "content_block_stop", "index": 0}\n\n', + 'event: content_block_start\ndata: {"type": "content_block_start", "index": 1, ' + '"content_block": {"type": "tool_use", "id": "toolu_synthetic", "name": "lookup_weather", "input": {}}}\n\n', + 'event: content_block_delta\ndata: {"type": "content_block_delta", "index": 1, ' + '"delta": {"type": "input_json_delta", "partial_json": "{\\"city\\": \\"Paris\\"}"}}\n\n', + 'event: content_block_stop\ndata: {"type": "content_block_stop", "index": 1}\n\n', + 'event: message_delta\ndata: {"type": "message_delta", "delta": {"stop_reason": "tool_use", ' + '"stop_sequence": null}, "usage": {"output_tokens": 9}}\n\n', + 'event: message_stop\ndata: {"type": "message_stop"}\n\n', + ) + tools: Final = [ + { + "name": "lookup_weather", + "description": "Look up the forecast for a city", + "input_schema": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]}, + } + ] + + def guardrail(request: Request) -> Reply: + assert request.target == f"/guardrail/{guardrail_id}/version/DRAFT/apply", request.target + body: Final = json.loads(request.body) + assert body["source"] == "OUTPUT", body + assert body["content"] == [{"text": {"text": spoken}}], body + return Reply(body=json.dumps({"action": "NONE", "outputs": [], "assessments": []}).encode()) + + def provider(request: Request) -> Reply: + assert request.target == "/v1/messages" + body: Final = json.loads(request.body) + assert body["stream"] is True, body + assert body["tools"] == tools, body + return Reply(content_type="text/event-stream", chunks=tuple(frame.encode() for frame in frames)) + + with wire_server(guardrail) as policy, wire_server(provider) as upstream: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["guardrails"] = [ + { + "guardrail_name": identity, + "litellm_params": { + "guardrail": "bedrock", + "mode": "post_call", + "default_on": True, + "mask_response_content": True, + "guardrailIdentifier": guardrail_id, + "guardrailVersion": "DRAFT", + "aws_region_name": "us-east-1", + "aws_access_key_id": "AKIASYNTHETICGUARDRAIL", + "aws_secret_access_key": "synthetic-secret", + "aws_bedrock_runtime_endpoint": policy.url, + }, + } + ] + path: Final = tmp_path / "bedrock-stream.yaml" + path.write_text(yaml.safe_dump(config)) + with owned_proxy(gateway, tmp_path, {}, config=path) as candidate, candidate.scenario() as scenario: + model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", api_base=upstream.url, api_key="synthetic-anthropic-key" + ) + response: Final = candidate.request( + "POST", + "/v1/messages", + { + "model": model, + "max_tokens": 64, + "stream": True, + "tools": tools, + "messages": [{"role": "user", "content": f"What is the weather in Paris? {identity}"}], + }, + ) + assert response.status_code == 200, response.text + head, separator, tail = response.text.partition("\n\n") + assert separator == "\n\n", response.text + assert head.startswith("event: message_start\ndata: "), response.text + assert json.loads(head.removeprefix("event: message_start\ndata: ")) == { + "type": "message_start", + "message": { + "id": "msg_synthetic", + "type": "message", + "role": "assistant", + "model": model, + "content": [], + "stop_reason": None, + "stop_sequence": None, + "usage": {"input_tokens": 11, "output_tokens": 1}, + }, + }, response.text + assert tail == "".join(frames[1:]), response.text + assert len(policy.drain()) == len(upstream.drain()) == 1 + + @pytest.mark.covers("other.mcp.guardrails.request_selection_blocks_resolved_tool_without_execution") def test_request_selected_mcp_guardrail_blocks_direct_and_virtual_calls(gateway: Gateway, tmp_path: Path) -> None: guardrail = "mcp-policy-" + uuid.uuid4().hex diff --git a/tests/integration/pricing/test_configured_prices.py b/tests/integration/pricing/test_configured_prices.py index e39833516c2..0e4efea3a15 100644 --- a/tests/integration/pricing/test_configured_prices.py +++ b/tests/integration/pricing/test_configured_prices.py @@ -8,8 +8,10 @@ import pytest import yaml from pydantic import JsonValue +from litellm import get_model_info from tests.integration._support.client import Gateway, eventually, object_value, string_value from tests.integration._support.database import read_rows +from tests.integration._support.process import owned_proxy @pytest.mark.covers("quota_management.spend_tracking.custom_price.matches_input_rates") @@ -169,6 +171,49 @@ def test_saving_echoed_model_info_does_not_freeze_cost_map_price_into_deployment assert {key: value for key, value in stored.items() if key in COST_MAP_DISPLAY_PRICING_KEYS} == {}, stored +def displayed_model_info(gateway: Gateway, model: str) -> dict[str, JsonValue]: + entries: Final = gateway.get("/model/info")["data"] + assert isinstance(entries, list) + target: Final = next(object_value(entry) for entry in entries if object_value(entry)["model_name"] == model) + return object_value(target["model_info"]) + + +@pytest.mark.covers("pricing.model_update.echoed_cost_map_metadata_is_not_persisted_as_override") +def test_saving_echoed_model_info_does_not_persist_cost_map_metadata_as_overrides(gateway: Gateway) -> None: + catalog_entry: Final = get_model_info("openai/gpt-4o-mini") + with gateway.scenario() as scenario: + model: Final = scenario.model() + displayed: Final = displayed_model_info(gateway, model) + identity: Final = string_value(displayed["id"]) + assert displayed["key"] == catalog_entry["key"], displayed + assert displayed["max_input_tokens"] == catalog_entry["max_input_tokens"], displayed + saved: Final = gateway.request( + "PATCH", f"/model/{identity}/update", {"model_info": {**displayed, "description": "echoed ui save"}} + ) + assert saved.status_code == 200, saved.text + stored: Final = persisted_model_info(identity) + assert stored["description"] == "echoed ui save", stored + assert {key: value for key, value in stored.items() if key in catalog_entry} == {}, stored + + +@pytest.mark.covers("pricing.model_update.echoing_cost_map_value_back_clears_stored_override") +def test_saving_the_cost_map_value_back_over_a_stored_override_clears_it(gateway: Gateway, tmp_path: Path) -> None: + catalog_limit: Final = get_model_info("openai/gpt-4o-mini")["max_input_tokens"] + assert isinstance(catalog_limit, int) and catalog_limit != 4321, catalog_limit + with owned_proxy(gateway, tmp_path, {}) as candidate, candidate.scenario() as scenario: + overridden: Final = scenario.model(model_info={"max_input_tokens": 4321}) + displayed: Final = displayed_model_info(candidate, overridden) + identity: Final = string_value(displayed["id"]) + assert displayed["max_input_tokens"] == 4321, displayed + assert persisted_model_info(identity)["max_input_tokens"] == 4321 + saved: Final = candidate.request( + "PATCH", f"/model/{identity}/update", {"model_info": {**displayed, "max_input_tokens": catalog_limit}} + ) + assert saved.status_code == 200, saved.text + stored: Final = persisted_model_info(identity) + assert "max_input_tokens" not in stored, stored + + @pytest.mark.covers("quota_management.spend_tracking.default_prices.loaded_router_preserves_cached_defaults") def test_loaded_router_preserves_cached_defaults_during_real_requests(gateway: Gateway, tmp_path: Path) -> None: from litellm import Router diff --git a/tests/integration/pricing/test_realtime_cached_audio_pricing.py b/tests/integration/pricing/test_realtime_cached_audio_pricing.py new file mode 100644 index 00000000000..4a7598d0cbf --- /dev/null +++ b/tests/integration/pricing/test_realtime_cached_audio_pricing.py @@ -0,0 +1,131 @@ +import asyncio +import json +import os +import uuid +from hashlib import sha256 +from typing import Final + +import pytest +import websockets +from pydantic import BaseModel, ConfigDict, JsonValue + +import litellm +from tests.integration._support.client import JSON_OBJECT, Gateway, eventually +from tests.integration._support.database import read_rows +from tests.integration._support.upstream import delete_scenario, register_scenario +from tests.integration.cost_calculation.cost_tracking_case import RealtimeResponse + + +class RealtimeRates(BaseModel): + model_config = ConfigDict(frozen=True) + + input_cost_per_token: float + input_cost_per_audio_token: float + cache_read_input_token_cost: float + cache_read_input_audio_token_cost: float | None = None + output_cost_per_token: float + output_cost_per_audio_token: float + + +MODEL: Final = "gpt-realtime-2" +RATES: Final = RealtimeRates.model_validate(litellm.get_model_info(MODEL, custom_llm_provider="openai")) +TEXT_RATE: Final = RATES.input_cost_per_token +AUDIO_RATE: Final = RATES.input_cost_per_audio_token +CACHED_TEXT_RATE: Final = RATES.cache_read_input_token_cost +CACHED_AUDIO_RATE: Final = ( + RATES.cache_read_input_token_cost + if RATES.cache_read_input_audio_token_cost is None + else RATES.cache_read_input_audio_token_cost +) +OUTPUT_TEXT_RATE: Final = RATES.output_cost_per_token +OUTPUT_AUDIO_RATE: Final = RATES.output_cost_per_audio_token +INPUT_TEXT_TOKENS: Final = 116 +INPUT_AUDIO_TOKENS: Final = 167 +CACHED_TEXT_TOKENS: Final = 64 +CACHED_AUDIO_TOKENS: Final = 128 +INPUT_TOKENS: Final = INPUT_TEXT_TOKENS + INPUT_AUDIO_TOKENS +OUTPUT_TEXT_TOKENS: Final = 8 +OUTPUT_AUDIO_TOKENS: Final = 12 +OUTPUT_TOKENS: Final = OUTPUT_TEXT_TOKENS + OUTPUT_AUDIO_TOKENS +EXPECTED_INPUT_COST: Final = ( + (INPUT_TEXT_TOKENS - CACHED_TEXT_TOKENS) * TEXT_RATE + + CACHED_TEXT_TOKENS * CACHED_TEXT_RATE + + (INPUT_AUDIO_TOKENS - CACHED_AUDIO_TOKENS) * AUDIO_RATE + + CACHED_AUDIO_TOKENS * CACHED_AUDIO_RATE +) +EXPECTED_OUTPUT_COST: Final = OUTPUT_TEXT_TOKENS * OUTPUT_TEXT_RATE + OUTPUT_AUDIO_TOKENS * OUTPUT_AUDIO_RATE + + +def cached_audio_response_done() -> RealtimeResponse: + return RealtimeResponse( + content_type="application/x-realtime", + events=( + { + "type": "response.done", + "event_id": "evt_$REQUEST_ID", + "response": { + "id": "resp_$REQUEST_ID", + "object": "realtime.response", + "status": "completed", + "output": [], + "usage": { + "total_tokens": INPUT_TOKENS + OUTPUT_TOKENS, + "input_tokens": INPUT_TOKENS, + "output_tokens": OUTPUT_TOKENS, + "input_token_details": { + "text_tokens": INPUT_TEXT_TOKENS, + "audio_tokens": INPUT_AUDIO_TOKENS, + "cached_tokens": CACHED_TEXT_TOKENS + CACHED_AUDIO_TOKENS, + "cached_tokens_details": { + "text_tokens": CACHED_TEXT_TOKENS, + "audio_tokens": CACHED_AUDIO_TOKENS, + }, + }, + "output_token_details": { + "text_tokens": OUTPUT_TEXT_TOKENS, + "audio_tokens": OUTPUT_AUDIO_TOKENS, + }, + }, + }, + }, + ), + ) + + +async def _one_realtime_turn(proxy_url: str, key: str, model: str) -> dict[str, JsonValue]: + async with websockets.connect( + f"{proxy_url.replace('http://', 'ws://').replace('https://', 'wss://')}/v1/realtime?model={model}", + additional_headers={"Authorization": f"Bearer {key}"}, + ) as websocket: + session: Final = JSON_OBJECT.validate_json(await websocket.recv()) + await websocket.send(json.dumps({"type": "response.create"})) + async for message in websocket: + if JSON_OBJECT.validate_json(message).get("type") == "response.done": + return session + raise AssertionError(f"websocket closed before response.done for {model}") + + +@pytest.mark.covers("pricing.realtime.cached_audio_tokens_bill_at_audio_cache_read_rate") +def test_realtime_cached_audio_tokens_bill_at_audio_cache_read_rate_not_full_audio_rate(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + scenario_id: Final = f"realtime-cached-audio-{uuid.uuid4().hex[:12]}" + handle: Final = register_scenario(scenario_id, cached_audio_response_done()) + scenario.cleanups.callback(delete_scenario, handle) + key: Final = scenario.key() + model: Final = scenario.model( + model=f"openai/{MODEL}", api_key=scenario_id, api_base=gateway.upstream_url.rstrip("/") + ) + session: Final = asyncio.run(_one_realtime_turn(os.environ["INTEGRATION_PROXY_URL"].rstrip("/"), key, model)) + assert session.get("type") == "session.created", session + rows: Final = eventually( + lambda: read_rows( + 'SELECT spend, prompt_tokens, completion_tokens, call_type FROM "LiteLLM_SpendLogs" WHERE api_key = %s', + (sha256(key.encode()).hexdigest(),), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert rows[0]["call_type"] == "_arealtime", rows + assert rows[0]["prompt_tokens"] == INPUT_TOKENS, rows + assert rows[0]["completion_tokens"] == OUTPUT_TOKENS, rows + assert float(str(rows[0]["spend"])) == pytest.approx(EXPECTED_INPUT_COST + EXPECTED_OUTPUT_COST, rel=1e-6), rows diff --git a/tests/integration/spend/test_cache_and_quota.py b/tests/integration/spend/test_cache_and_quota.py index 97ef785aa15..1c2b5551855 100644 --- a/tests/integration/spend/test_cache_and_quota.py +++ b/tests/integration/spend/test_cache_and_quota.py @@ -1,19 +1,27 @@ import json +import os import threading import uuid +from collections.abc import Generator from concurrent.futures import ThreadPoolExecutor -from contextlib import ExitStack +from contextlib import ExitStack, contextmanager from hashlib import sha256 +from pathlib import Path from typing import Final +from urllib.parse import urlsplit, urlunsplit import httpx +import psycopg import pytest from hypothesis import strategies as st from hypothesis.stateful import RuleBasedStateMachine, rule, run_state_machine_as_test -from integration._support.client import Gateway, eventually +from integration._support.client import Gateway, eventually, string_value from integration._support.database import read_rows +from integration._support.database_relay import database_relay from integration._support.generation import LIFECYCLE_SETTINGS, bounded_http_requests +from integration._support.process import owned_proxy from integration._support.wire import Reply, Request, wire_server +from psycopg import sql @pytest.mark.covers("quota_management.response_cache.generated_sequences_preserve_content_and_accounting") @@ -214,6 +222,99 @@ def test_key_budget_at_boundary_blocks_provider_then_explicit_reset_restores(gat assert upstream.get("/__observations").json()["requests"] == [] +RESET_SWEEP_QUERY: Final = b'"LiteLLM_VerificationToken"."budget_reset_at" < $' + + +@contextmanager +def scratch_database() -> Generator[str]: + name: Final = f"integration_{uuid.uuid4().hex}" + with psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as admin: + admin.execute(sql.SQL("CREATE DATABASE {}").format(sql.Identifier(name))) + try: + yield urlunsplit(urlsplit(os.environ["DATABASE_URL"])._replace(path=f"/{name}")) + finally: + admin.execute(sql.SQL("DROP DATABASE {} WITH (FORCE)").format(sql.Identifier(name))) + + +@pytest.mark.covers("quota_management.budget.key.scheduled_reset_survives_transient_db_outage") +@pytest.mark.timeout(300) +def test_scheduled_budget_reset_reconnects_after_db_transport_failure_and_unblocks_key( + gateway: Gateway, tmp_path: Path +) -> None: + with ( + scratch_database() as scratch_url, + database_relay(scratch_url, RESET_SWEEP_QUERY) as (relay, relayed_url), + httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream, + owned_proxy( + gateway, + tmp_path, + { + "DATABASE_URL": relayed_url, + "PROXY_BUDGET_RESCHEDULER_MIN_TIME": "30", + "PROXY_BUDGET_RESCHEDULER_MAX_TIME": "30", + "PRISMA_HEALTH_WATCHDOG_ENABLED": "false", + }, + ) as candidate, + ): + model: Final = f"integration-{uuid.uuid4().hex}" + candidate.post( + "/model/new", + { + "model_name": model, + "litellm_params": { + "model": "openai/gpt-4o-mini", + "api_key": "integration-provider-key", + "api_base": f"{gateway.upstream_url}/v1", + "input_cost_per_token": 0.001, + "output_cost_per_token": 0.002, + }, + "model_info": {}, + }, + ) + key: Final = string_value( + candidate.post("/key/generate", {"models": [model], "max_budget": 0.06, "budget_duration": "5s"})["key"] + ) + digest: Final = sha256(key.encode()).hexdigest() + row_query: Final = ( + 'SELECT spend, budget_reset_at::text AS budget_reset_at FROM "LiteLLM_VerificationToken" WHERE token=%s' + ) + assert candidate.chat(model, key=key, text=f"spend it {uuid.uuid4().hex}")["usage"]["total_tokens"] == 40 + exhausted: Final = eventually( + lambda: read_rows(row_query, (digest,), database_url=scratch_url), + lambda rows: len(rows) == 1 and float(rows[0]["spend"]) >= 0.06, + seconds=70, + ) + assert float(exhausted[0]["spend"]) == pytest.approx(0.06) + upstream.get("/__observations").raise_for_status() + denied: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": f"over budget {uuid.uuid4().hex}"}]}, + key=key, + ) + assert denied.status_code == 422 and denied.json()["error"]["type"] == "budget_exceeded", denied.text + assert upstream.get("/__observations").json()["requests"] == [] + relay.arm() + assert relay.tripped.wait(90), "Scheduled reset sweep never reached the database" + eventually(lambda: relay.refused, lambda count: count >= 1, seconds=30) + reset: Final = eventually( + lambda: read_rows(row_query, (digest,), database_url=scratch_url), + lambda rows: len(rows) == 1 and float(rows[0]["spend"]) == 0, + seconds=80, + return_last_on_timeout=True, + ) + assert len(reset) == 1 and reset[0]["spend"] == 0.0, (exhausted, reset) + assert str(reset[0]["budget_reset_at"]) > str(exhausted[0]["budget_reset_at"]), (exhausted, reset) + prompt: Final = f"after reset {uuid.uuid4().hex}" + recovered: Final = candidate.request( + "POST", "/v1/chat/completions", {"model": model, "messages": [{"role": "user", "content": prompt}]}, key=key + ) + assert recovered.status_code == 200, recovered.text + assert recovered.json()["usage"]["total_tokens"] == 40, recovered.text + reached: Final = upstream.get("/__observations").json()["requests"] + assert len(reached) == 1 and reached[0]["body"]["messages"] == [{"role": "user", "content": prompt}], reached + + @pytest.mark.covers("quota_management.budget.key.count_tokens_reserves_nothing_so_completion_within_budget_succeeds") def test_repeated_count_tokens_on_budgeted_key_does_not_reserve_budget_or_block_later_completion( gateway: Gateway, @@ -333,6 +434,7 @@ def test_different_system_messages_do_not_share_a_cached_response(gateway: Gatew ): model: Final = scenario.model() prompt: Final = uuid.uuid4().hex + def completion_id(system: str, expected_calls: int) -> str: upstream.get("/__observations").raise_for_status() response: Final = gateway.request( diff --git a/tests/integration/spend/test_end_user_spend_without_proxy_user.py b/tests/integration/spend/test_end_user_spend_without_proxy_user.py new file mode 100644 index 00000000000..15d77f16db4 --- /dev/null +++ b/tests/integration/spend/test_end_user_spend_without_proxy_user.py @@ -0,0 +1,29 @@ +import uuid +from typing import Final + +import pytest +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows + + +@pytest.mark.covers("spend.end_user.charged_when_key_has_no_user_id_and_auth_cache_is_redis") +def test_end_user_spend_lands_for_key_without_user_id_when_auth_cache_is_redis(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002) + key: Final = scenario.key(models=[model]) + end_user: Final = f"integration-end-user-{uuid.uuid4().hex}" + scenario.cleanups.callback(gateway.request, "POST", "/customer/delete", {"user_ids": [end_user]}) + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": f"end user spend {end_user}"}], "user": end_user}, + key=key, + ) + assert response.status_code == 200, response.text + assert response.json()["usage"]["total_tokens"] == 40, response.text + charged: Final = eventually( + lambda: read_rows('SELECT spend FROM "LiteLLM_EndUserTable" WHERE user_id=%s', (end_user,)), + lambda values: len(values) == 1 and float(values[0]["spend"]) >= 0.06, + seconds=70, + ) + assert float(charged[0]["spend"]) == pytest.approx(20 * 0.001 + 20 * 0.002) diff --git a/tests/integration/spend/test_legacy_spend_logs_row_cap.py b/tests/integration/spend/test_legacy_spend_logs_row_cap.py new file mode 100644 index 00000000000..3a7ac8a83bd --- /dev/null +++ b/tests/integration/spend/test_legacy_spend_logs_row_cap.py @@ -0,0 +1,46 @@ +import os +import uuid +from datetime import datetime, timedelta, timezone +from typing import Final + +import psycopg +import pytest +from integration._support.client import Gateway +from integration._support.database import read_rows + +LEGACY_SPEND_LOGS_ROW_CAP: Final = 10000 + + +def _seed_spend_rows(user_id: str, count: int) -> tuple[str, ...]: + started: Final = datetime.now(timezone.utc) - timedelta(days=1) + with psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as connection: + connection.execute( + 'INSERT INTO "LiteLLM_SpendLogs" ' + '(request_id, call_type, "startTime", "endTime", "user", status) ' + "SELECT %s || '-' || n, 'acompletion', %s + n * interval '1 second', %s + n * interval '1 second', %s, " + "'success' FROM generate_series(1, %s) AS n", + (user_id, started, started, user_id, count), + ) + return tuple(f"{user_id}-{n}" for n in range(1, count + 1)) + + +def _delete_spend_rows(user_id: str) -> None: + with psycopg.connect(os.environ["DATABASE_URL"], autocommit=True) as connection: + connection.execute('DELETE FROM "LiteLLM_SpendLogs" WHERE "user" = %s', (user_id,)) + + +@pytest.mark.covers("spend.legacy_spend_logs.row_count_is_capped_at_the_most_recent_rows_and_flagged_truncated") +def test_legacy_spend_logs_returns_only_the_cap_of_most_recent_rows_and_flags_truncation(gateway: Gateway) -> None: + user_id: Final = f"integration-cap-{uuid.uuid4().hex}" + with gateway.scenario() as scenario: + scenario.cleanups.callback(_delete_spend_rows, user_id) + seeded: Final = _seed_spend_rows(user_id, LEGACY_SPEND_LOGS_ROW_CAP + 1) + assert read_rows('SELECT count(*)::text AS total FROM "LiteLLM_SpendLogs" WHERE "user" = %s', (user_id,)) == [ + {"total": str(LEGACY_SPEND_LOGS_ROW_CAP + 1)} + ] + response: Final = gateway.request("GET", "/spend/logs", params={"user_id": user_id}) + assert response.status_code == 200, response.text + returned: Final = tuple(row["request_id"] for row in response.json()) + assert len(returned) == LEGACY_SPEND_LOGS_ROW_CAP, f"{len(returned)} rows: {response.text[:300]}" + assert returned == tuple(reversed(seeded[1:])), response.text[:300] + assert response.headers.get("x-litellm-spend-logs-truncated") == "true", dict(response.headers) diff --git a/tests/integration/spend/test_messages_stream_usage_cost.py b/tests/integration/spend/test_messages_stream_usage_cost.py new file mode 100644 index 00000000000..692dd01e0b0 --- /dev/null +++ b/tests/integration/spend/test_messages_stream_usage_cost.py @@ -0,0 +1,176 @@ +import base64 +import json +import uuid +from pathlib import Path +from typing import Final + +import pytest +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.process import owned_proxy +from integration._support.upstream import _aws_event_frame +from integration._support.wire import Reply, Request, wire_server +from pydantic import JsonValue, TypeAdapter + +BEDROCK_MODEL: Final = "anthropic.claude-haiku-4-5-20251001-v1:0" +INPUT_TOKENS: Final = 30 +CACHE_READ_TOKENS: Final = 900 +CACHE_CREATION_TOKENS: Final = 400 +OUTPUT_TOKENS: Final = 57 +INPUT_RATE: Final = 0.001 +OUTPUT_RATE: Final = 0.002 +CACHE_READ_RATE: Final = 0.0001 +CACHE_CREATION_RATE: Final = 0.00125 +STREAMED_USAGE: Final = TypeAdapter(dict[str, float]) +EXPECTED_SPEND: Final = ( + INPUT_TOKENS * INPUT_RATE + + CACHE_READ_TOKENS * CACHE_READ_RATE + + CACHE_CREATION_TOKENS * CACHE_CREATION_RATE + + OUTPUT_TOKENS * OUTPUT_RATE +) + + +def _invoke_chunk(payload: dict[str, JsonValue]) -> bytes: + encoded: Final = base64.b64encode(json.dumps(payload, separators=(",", ":")).encode()).decode() + return _aws_event_frame("chunk", {"bytes": encoded}, "", "") + + +def _stream(message_id: str) -> bytes: + return ( + _invoke_chunk( + { + "type": "message_start", + "message": { + "id": message_id, + "type": "message", + "role": "assistant", + "model": BEDROCK_MODEL, + "content": [], + "stop_reason": None, + "stop_sequence": None, + "usage": { + "input_tokens": INPUT_TOKENS, + "cache_read_input_tokens": CACHE_READ_TOKENS, + "cache_creation_input_tokens": CACHE_CREATION_TOKENS, + "output_tokens": 0, + }, + }, + } + ) + + _invoke_chunk({"type": "content_block_start", "index": 0, "content_block": {"type": "text", "text": ""}}) + + _invoke_chunk( + {"type": "content_block_delta", "index": 0, "delta": {"type": "text_delta", "text": "cached answer"}} + ) + + _invoke_chunk({"type": "content_block_stop", "index": 0}) + + _invoke_chunk( + { + "type": "message_delta", + "delta": {"stop_reason": "end_turn"}, + "usage": { + "input_tokens": INPUT_TOKENS, + "cache_read_input_tokens": CACHE_READ_TOKENS, + "cache_creation_input_tokens": CACHE_CREATION_TOKENS, + "output_tokens": OUTPUT_TOKENS, + }, + } + ) + + _invoke_chunk({"type": "message_stop"}) + ) + + +def _proxy_config(directory: Path, model: str, upstream_url: str) -> Path: + config: Final = directory / "streamed_usage_cost_config.yaml" + config.write_text( + json.dumps( + { + "model_list": [ + { + "model_name": model, + "litellm_params": { + "model": f"bedrock/invoke/{BEDROCK_MODEL}", + "api_base": upstream_url, + "aws_access_key_id": "AKIASCRIPTEDPROVIDER", + "aws_secret_access_key": "scripted-secret", + "aws_region_name": "us-east-1", + "input_cost_per_token": INPUT_RATE, + "output_cost_per_token": OUTPUT_RATE, + "cache_read_input_token_cost": CACHE_READ_RATE, + "cache_creation_input_token_cost": CACHE_CREATION_RATE, + }, + } + ], + "general_settings": { + "master_key": "os.environ/LITELLM_MASTER_KEY", + "database_url": "os.environ/DATABASE_URL", + "store_model_in_db": True, + "disable_spend_logs": False, + "proxy_batch_write_at": 1, + "proxy_batch_polling_interval": 1, + }, + "litellm_settings": {"include_cost_in_streaming_usage": True}, + "router_settings": {"disable_cooldowns": True}, + } + ) + ) + return config + + +def _data_events(body: str) -> tuple[dict[str, JsonValue], ...]: + return tuple(json.loads(line.removeprefix("data:")) for line in body.splitlines() if line.startswith("data:")) + + +@pytest.mark.covers("spend.anthropic_messages_stream.streamed_usage_cost_equals_recorded_spend") +@pytest.mark.timeout(180) +def test_bedrock_messages_stream_usage_cost_matches_recorded_spend_with_custom_cache_rates( + gateway: Gateway, tmp_path: Path +) -> None: + message_id: Final = f"msg_{uuid.uuid4().hex}" + model: Final = f"{BEDROCK_MODEL}-{uuid.uuid4().hex}" + + def respond(request: Request) -> Reply: + assert request.target == f"/model/{BEDROCK_MODEL}/invoke-with-response-stream", request.target + assert json.loads(request.body)["messages"] == [{"role": "user", "content": "cached cost control"}], ( + request.body + ) + return Reply(content_type="application/vnd.amazon.eventstream", chunks=(_stream(message_id),)) + + with wire_server(respond) as wire: + config: Final = _proxy_config(tmp_path, model, wire.url) + with owned_proxy(gateway, tmp_path, {}, config=config) as candidate: + response: Final = candidate.request( + "POST", + "/v1/messages", + { + "model": model, + "messages": [{"role": "user", "content": "cached cost control"}], + "max_tokens": OUTPUT_TOKENS, + "stream": True, + }, + ) + assert response.status_code == 200, response.text + message_delta: Final = next( + event for event in _data_events(response.text) if event["type"] == "message_delta" + ) + streamed_usage: Final = STREAMED_USAGE.validate_python(message_delta["usage"]) + rows: Final = eventually( + lambda: read_rows( + 'SELECT status, prompt_tokens, completion_tokens, spend FROM "LiteLLM_SpendLogs" ' + "WHERE request_id=%s", + (message_id,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert rows[0]["status"] == "success", rows + assert rows[0]["prompt_tokens"] == INPUT_TOKENS + CACHE_READ_TOKENS + CACHE_CREATION_TOKENS, rows + assert rows[0]["completion_tokens"] == OUTPUT_TOKENS, rows + recorded_spend: Final = float(str(rows[0]["spend"])) + assert recorded_spend == pytest.approx(EXPECTED_SPEND), rows + assert streamed_usage == { + "input_tokens": INPUT_TOKENS, + "cache_read_input_tokens": CACHE_READ_TOKENS, + "cache_creation_input_tokens": CACHE_CREATION_TOKENS, + "output_tokens": OUTPUT_TOKENS, + "cost": pytest.approx(recorded_spend), + }, (streamed_usage, rows, response.text) + assert len(wire.drain()) == 1 diff --git a/tests/integration/spend/test_reset_budget_leader_election.py b/tests/integration/spend/test_reset_budget_leader_election.py new file mode 100644 index 00000000000..f9c165a0be0 --- /dev/null +++ b/tests/integration/spend/test_reset_budget_leader_election.py @@ -0,0 +1,75 @@ +import json +import os +from pathlib import Path +from types import MappingProxyType +from typing import Final + +import psycopg +import pytest +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.process import owned_proxy +from redis import Redis + +RESET_LEASE_KEY: Final = "cronjob_lock:reset_budget_job" +PEER_POD_LEASE: Final = json.dumps("integration-peer-pod-holding-the-reset-lease") +FAST_RESET_TICK: Final = MappingProxyType( + {"PROXY_BUDGET_RESCHEDULER_MIN_TIME": "2", "PROXY_BUDGET_RESCHEDULER_MAX_TIME": "2"} +) + + +def _team_spend(team: str) -> float: + rows: Final = read_rows('SELECT spend FROM "LiteLLM_TeamTable" WHERE team_id = %s', (team,)) + assert len(rows) == 1, rows + spend: Final = rows[0]["spend"] + assert isinstance(spend, (int, float)), rows + return float(spend) + + +def _make_team_budget_due(team: str, spend: float) -> None: + with psycopg.connect(os.environ["DATABASE_URL"]) as connection: + connection.execute( + "UPDATE \"LiteLLM_TeamTable\" SET spend = %s, budget_reset_at = now() - interval '1 day' " + "WHERE team_id = %s", + (spend, team), + ) + + +@pytest.mark.covers("spend.budget_reset.one_pod_sweeps_per_tick") +def test_reset_sweep_skips_ticks_while_another_pod_holds_the_lease_and_resumes_after_release( + gateway: Gateway, tmp_path: Path +) -> None: + with ( + gateway.scenario() as scenario, + Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) as cache, + ): + model: Final = scenario.model(input_cost_per_token=0.0, output_cost_per_token=0.0) + team: Final = scenario.team(max_budget=1.0, budget_duration="30d") + key: Final = scenario.key(team_id=team) + _make_team_budget_due(team, spend=0.5) + assert _team_spend(team) == 0.5 + eventually( + lambda: cache.set(RESET_LEASE_KEY, PEER_POD_LEASE, ex=120, nx=True), + lambda claimed: claimed is True, + seconds=60, + ) + try: + with owned_proxy(gateway, tmp_path, FAST_RESET_TICK) as replica: + response: Final = replica.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": "classification request"}]}, + key=key, + ) + assert response.status_code == 200, response.text + held: Final = eventually( + lambda: _team_spend(team), lambda spend: spend != 0.5, seconds=8, return_last_on_timeout=True + ) + assert held == 0.5, f"team {team} was swept while another pod held the reset lease: spend={held}" + assert cache.get(RESET_LEASE_KEY) == PEER_POD_LEASE.encode() + cache.delete(RESET_LEASE_KEY) + swept: Final = eventually(lambda: _team_spend(team), lambda spend: spend == 0.0, seconds=15) + assert swept == 0.0 + eventually(lambda: cache.get(RESET_LEASE_KEY), lambda value: value is None, seconds=15) + finally: + cache.delete(RESET_LEASE_KEY) diff --git a/tests/integration/spend/test_responses_cache_write_itemization.py b/tests/integration/spend/test_responses_cache_write_itemization.py new file mode 100644 index 00000000000..19d73b0aacf --- /dev/null +++ b/tests/integration/spend/test_responses_cache_write_itemization.py @@ -0,0 +1,98 @@ +import json +import uuid +from typing import Final + +import pytest + +from tests.integration._support.client import Gateway, eventually, object_value, string_value +from tests.integration._support.database import read_rows +from tests.integration._support.upstream import delete_scenario, register_scenario +from tests.integration.cost_calculation.cost_tracking_case import JsonResponse + +INPUT_RATE: Final = 0.001 +OUTPUT_RATE: Final = 0.002 +CACHE_CREATION_RATE: Final = 0.004 +CACHE_READ_RATE: Final = 0.0001 +UNCACHED_INPUT_TOKENS: Final = 1000 +CACHE_WRITE_TOKENS: Final = 2000 +CACHED_TOKENS: Final = 8000 +INPUT_TOKENS: Final = UNCACHED_INPUT_TOKENS + CACHE_WRITE_TOKENS + CACHED_TOKENS +OUTPUT_TOKENS: Final = 500 + + +def responses_cache_write_response() -> JsonResponse: + return JsonResponse( + content_type="application/json", + body={ + "id": "resp_$REQUEST_ID", + "object": "response", + "created_at": 1700000000, + "status": "completed", + "model": "gpt-5.6", + "output": [ + { + "type": "message", + "id": "msg_$REQUEST_ID", + "status": "completed", + "role": "assistant", + "content": [{"type": "output_text", "text": "scripted response", "annotations": []}], + } + ], + "usage": { + "input_tokens": INPUT_TOKENS, + "output_tokens": OUTPUT_TOKENS, + "total_tokens": INPUT_TOKENS + OUTPUT_TOKENS, + "input_tokens_details": { + "cached_tokens": CACHED_TOKENS, + "cache_write_tokens": CACHE_WRITE_TOKENS, + }, + "output_tokens_details": {"reasoning_tokens": 0}, + }, + }, + ) + + +@pytest.mark.covers("spend.responses_api.cache_write_tokens_itemized_as_cache_creation_cost") +def test_responses_cache_write_tokens_are_itemized_as_cache_creation_cost(gateway: Gateway) -> None: + with gateway.scenario() as scenario: + scenario_id: Final = f"responses-cache-write-{uuid.uuid4().hex[:12]}" + handle: Final = register_scenario(scenario_id, responses_cache_write_response()) + scenario.cleanups.callback(delete_scenario, handle) + model: Final = scenario.model( + model="openai/gpt-5.6", + api_base=handle.api_base(), + input_cost_per_token=INPUT_RATE, + output_cost_per_token=OUTPUT_RATE, + cache_creation_input_token_cost=CACHE_CREATION_RATE, + cache_read_input_token_cost=CACHE_READ_RATE, + ) + response: Final = gateway.request("POST", "/v1/responses", {"model": model, "input": "cache write control"}) + assert response.status_code == 200, response.text + expected_cache_creation_cost: Final = CACHE_WRITE_TOKENS * CACHE_CREATION_RATE + expected_cache_read_cost: Final = CACHED_TOKENS * CACHE_READ_RATE + expected_input_cost: Final = ( + UNCACHED_INPUT_TOKENS * INPUT_RATE + expected_cache_creation_cost + expected_cache_read_cost + ) + expected_output_cost: Final = OUTPUT_TOKENS * OUTPUT_RATE + request_id: Final = string_value(object_value(response.json())["id"]) + rows: Final = eventually( + lambda: read_rows( + 'SELECT spend, metadata, prompt_tokens, completion_tokens FROM "LiteLLM_SpendLogs" ' + "WHERE request_id = %s", + (request_id,), + ), + lambda values: len(values) == 1, + seconds=70, + ) + assert rows[0]["prompt_tokens"] == INPUT_TOKENS, response.text + assert rows[0]["completion_tokens"] == OUTPUT_TOKENS, response.text + assert float(rows[0]["spend"]) == pytest.approx(expected_input_cost + expected_output_cost, rel=1e-6), ( + response.text + ) + metadata: Final = rows[0]["metadata"] + parsed: Final = json.loads(metadata) if isinstance(metadata, str) else object_value(metadata) + breakdown: Final = object_value(parsed["cost_breakdown"]) + assert breakdown.get("cache_creation_cost") == pytest.approx(expected_cache_creation_cost, rel=1e-6), breakdown + assert breakdown.get("cache_read_cost") == pytest.approx(expected_cache_read_cost, rel=1e-6), breakdown + assert float(breakdown["input_cost"]) == pytest.approx(expected_input_cost, rel=1e-6), breakdown + assert float(breakdown["output_cost"]) == pytest.approx(expected_output_cost, rel=1e-6), breakdown diff --git a/tests/integration/spend/test_session_total_spend.py b/tests/integration/spend/test_session_total_spend.py new file mode 100644 index 00000000000..54df0c03d81 --- /dev/null +++ b/tests/integration/spend/test_session_total_spend.py @@ -0,0 +1,98 @@ +import json +import uuid +from datetime import datetime, timedelta, timezone +from typing import Final + +import pytest +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.wire import Reply, Request, wire_server + +INPUT_COST_PER_TOKEN: Final = 0.001 +OUTPUT_COST_PER_TOKEN: Final = 0.002 +ROUND_USAGE: Final = ((10, 5), (20, 10), (30, 15)) +ROUND_SPEND: Final = tuple( + prompt * INPUT_COST_PER_TOKEN + completion * OUTPUT_COST_PER_TOKEN for prompt, completion in ROUND_USAGE +) +SESSION_SPEND: Final = sum(ROUND_SPEND) + + +def _round_reply(request: Request, prompt: str, usage: tuple[int, int]) -> Reply: + assert request.method == "POST" and request.target == "/chat/completions", request.target + assert json.loads(request.body) == { + "model": "gpt-4o-mini", + "messages": [{"role": "user", "content": prompt}], + }, request.body + return Reply( + body=json.dumps( + { + "id": "chatcmpl-" + uuid.uuid4().hex, + "object": "chat.completion", + "created": 1, + "model": "gpt-4o-mini", + "choices": [ + {"index": 0, "message": {"role": "assistant", "content": "round answer"}, "finish_reason": "stop"} + ], + "usage": {"prompt_tokens": usage[0], "completion_tokens": usage[1], "total_tokens": sum(usage)}, + } + ).encode() + ) + + +@pytest.mark.covers("spend.logs_ui.multi_round_session_total_spend_sums_every_round") +def test_logs_ui_session_total_spend_sums_every_round_of_a_multi_round_session(gateway: Gateway) -> None: + session_id: Final = f"session-{uuid.uuid4().hex}" + prompts: Final = tuple(f"round-{index}-{uuid.uuid4().hex}" for index in range(len(ROUND_USAGE))) + usage_by_prompt: Final = dict(zip(prompts, ROUND_USAGE, strict=True)) + + def provider(request: Request) -> Reply: + prompt: Final = str(json.loads(request.body)["messages"][0]["content"]) + return _round_reply(request, prompt, usage_by_prompt[prompt]) + + with wire_server(provider) as wire, gateway.scenario() as scenario: + alias: Final = scenario.model( + api_base=wire.url, + input_cost_per_token=INPUT_COST_PER_TOKEN, + output_cost_per_token=OUTPUT_COST_PER_TOKEN, + num_retries=0, + ) + request_ids: Final = tuple(_completed_round(gateway, alias, session_id, prompt) for prompt in prompts) + assert len(wire.drain()) == len(ROUND_USAGE) + rows: Final = eventually( + lambda: read_rows( + 'SELECT request_id, spend FROM "LiteLLM_SpendLogs" WHERE session_id=%s ORDER BY "startTime"', + (session_id,), + ), + lambda values: len(values) == len(ROUND_USAGE), + seconds=70, + ) + assert [row["request_id"] for row in rows] == list(request_ids), rows + assert [float(row["spend"]) for row in rows] == pytest.approx(list(ROUND_SPEND)), rows + now: Final = datetime.now(timezone.utc) + logs: Final = gateway.request( + "GET", + "/spend/logs/ui", + params={ + "session_id": session_id, + "start_date": (now - timedelta(days=1)).strftime("%Y-%m-%d %H:%M:%S"), + "end_date": (now + timedelta(days=1)).strftime("%Y-%m-%d %H:%M:%S"), + }, + ) + assert logs.status_code == 200, logs.text + page: Final = logs.json() + assert page["total"] == len(ROUND_USAGE), logs.text + assert sorted(row["request_id"] for row in page["data"]) == sorted(request_ids), logs.text + assert [row["session_total_count"] for row in page["data"]] == [len(ROUND_USAGE)] * len(ROUND_USAGE), logs.text + assert [row["session_total_spend"] for row in page["data"]] == pytest.approx( + [SESSION_SPEND] * len(ROUND_USAGE) + ), logs.text + + +def _completed_round(gateway: Gateway, alias: str, session_id: str, prompt: str) -> str: + response: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": alias, "messages": [{"role": "user", "content": prompt}], "litellm_trace_id": session_id}, + ) + assert response.status_code == 200, response.text + return str(response.json()["id"]) diff --git a/tests/integration/spend/test_spend_log_write_batching.py b/tests/integration/spend/test_spend_log_write_batching.py new file mode 100644 index 00000000000..386f99b1d07 --- /dev/null +++ b/tests/integration/spend/test_spend_log_write_batching.py @@ -0,0 +1,75 @@ +import uuid +from pathlib import Path +from typing import Final + +import pytest +import yaml +from integration._support.client import Gateway, eventually, string_value +from integration._support.database import read_rows +from integration._support.process import owned_proxy +from pydantic import JsonValue + +WRITE_STATEMENT_MAX_BYTES: Final = 200_000 +MESSAGES_PER_REQUEST: Final = 60 +MESSAGE_CHARACTERS: Final = 2_000 +REQUESTS: Final = 4 + + +def _config_storing_prompts(tmp_path: Path) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["general_settings"]["store_prompts_in_spend_logs"] = True + path: Final = tmp_path / "store-prompts.yaml" + path.write_text(yaml.safe_dump(config)) + return path + + +def _prompt_messages(marker: str) -> list[JsonValue]: + return [ + {"role": "user", "content": f"{marker}-{index}-".ljust(MESSAGE_CHARACTERS, "x")} + for index in range(MESSAGES_PER_REQUEST) + ] + + +def _persisted(request_ids: tuple[str, ...]) -> list[dict[str, JsonValue]]: + placeholders: Final = ", ".join("%s" for _ in request_ids) + return read_rows( + "SELECT request_id, xmin::text AS statement, octet_length(proxy_server_request::text) AS stored_bytes " + f'FROM "LiteLLM_SpendLogs" WHERE request_id IN ({placeholders}) ORDER BY request_id', + request_ids, + ) + + +@pytest.mark.covers("quota_management.spend_tracking.prompt_rows_are_written_in_byte_bounded_statements") +def test_prompt_carrying_spend_rows_flushed_together_are_written_in_byte_bounded_statements( + gateway: Gateway, tmp_path: Path +) -> None: + with ( + gateway.scenario() as scenario, + owned_proxy( + gateway, + tmp_path, + { + "SPEND_LOG_WRITE_BATCH_MAX_BYTES": str(WRITE_STATEMENT_MAX_BYTES), + "SPEND_LOG_QUEUE_POLL_INTERVAL": "15", + }, + config=_config_storing_prompts(tmp_path), + ) as owned, + ): + model: Final = scenario.model() + request_ids: Final = tuple( + string_value( + owned.post( + "/v1/chat/completions", + {"model": model, "messages": _prompt_messages(f"integration-prompt-{uuid.uuid4().hex}")}, + )["id"] + ) + for _ in range(REQUESTS) + ) + assert len(set(request_ids)) == REQUESTS, request_ids + rows: Final = eventually(lambda: _persisted(request_ids), lambda values: len(values) == REQUESTS, seconds=70) + stored_bytes: Final = tuple(row["stored_bytes"] for row in rows) + assert all( + isinstance(size, int) and WRITE_STATEMENT_MAX_BYTES // 2 < size < WRITE_STATEMENT_MAX_BYTES + for size in stored_bytes + ), rows + assert len({row["statement"] for row in rows}) == REQUESTS, rows diff --git a/tests/integration/spend/test_team_member_spend.py b/tests/integration/spend/test_team_member_spend.py index 89eb2a24bcf..cf1dac25793 100644 --- a/tests/integration/spend/test_team_member_spend.py +++ b/tests/integration/spend/test_team_member_spend.py @@ -1,9 +1,12 @@ +import os import uuid from typing import Final +import httpx import pytest from integration._support.client import Gateway, eventually, object_value from integration._support.database import read_rows +from redis import Redis @pytest.mark.covers("spend.team_member.member_without_budget_gets_membership_row_and_spend") @@ -52,3 +55,45 @@ def test_member_added_without_any_budget_is_charged_on_its_membership_row(gatewa if object_value(row)["user_id"] == user ] assert len(exposed) == 1 and exposed[0][1] == pytest.approx(0.06), info + + +@pytest.mark.covers("spend.team_member.stale_low_redis_counter_still_blocks_member_over_budget") +def test_member_over_budget_is_blocked_when_redis_counter_reads_stale_low(gateway: Gateway) -> None: + with ( + gateway.scenario() as scenario, + httpx.Client(base_url=gateway.upstream_url, timeout=5, trust_env=False) as upstream, + Redis(host=os.environ["REDIS_HOST"], port=int(os.environ["REDIS_PORT"])) as cache, + ): + model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002) + team: Final = scenario.team(models=[model]) + user: Final = scenario.user() + added: Final = gateway.request( + "POST", + "/team/member_add", + {"team_id": team, "member": {"user_id": user, "role": "user"}, "max_budget_in_team": 0.05}, + ) + assert added.status_code == 200, added.text + key: Final = scenario.key(team_id=team, user_id=user, models=[model]) + assert gateway.chat(model, key=key, text=f"member budget {uuid.uuid4().hex}")["usage"]["total_tokens"] == 40 + charged: Final = eventually( + lambda: read_rows( + 'SELECT spend FROM "LiteLLM_TeamMembership" WHERE team_id=%s AND user_id=%s', (team, user) + ), + lambda values: len(values) == 1 and float(values[0]["spend"]) >= 0.06, + seconds=70, + ) + assert float(charged[0]["spend"]) == pytest.approx(0.06) + counter_key: Final = f"spend:team_member:{user}:{team}" + counted: Final = eventually(lambda: cache.get(counter_key), lambda value: value is not None, seconds=10) + assert float(counted) == pytest.approx(0.06), counted + cache.set(counter_key, "0.01") + upstream.get("/__observations").raise_for_status() + denied: Final = gateway.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": f"stale counter {uuid.uuid4().hex}"}]}, + key=key, + ) + assert denied.status_code == 422 and denied.json()["error"]["type"] == "budget_exceeded", denied.text + assert upstream.get("/__observations").json()["requests"] == [] + assert float(cache.get(counter_key)) == pytest.approx(0.06), denied.text diff --git a/tests/integration/spend/test_user_budget_on_team_keys.py b/tests/integration/spend/test_user_budget_on_team_keys.py new file mode 100644 index 00000000000..f8e0eea3454 --- /dev/null +++ b/tests/integration/spend/test_user_budget_on_team_keys.py @@ -0,0 +1,54 @@ +import uuid +from hashlib import sha256 +from pathlib import Path +from typing import Final + +import httpx +import pytest +import yaml +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.process import owned_proxy + + +@pytest.mark.covers("spend.user_budget.opted_in_team_key_is_denied_once_owner_budget_is_exhausted") +def test_team_key_is_denied_before_provider_once_owner_personal_budget_is_exhausted_when_opted_in( + gateway: Gateway, tmp_path: Path +) -> None: + configuration: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + configuration["general_settings"]["apply_user_budget_to_team_keys"] = True + path: Final = tmp_path / "apply-user-budget-to-team-keys.yaml" + path.write_text(yaml.safe_dump(configuration)) + with ( + owned_proxy(gateway, tmp_path, {}, config=path) as candidate, + candidate.scenario() as scenario, + httpx.Client(base_url=candidate.upstream_url, timeout=5, trust_env=False) as upstream, + ): + model: Final = scenario.model(input_cost_per_token=0.001, output_cost_per_token=0.002) + user: Final = scenario.user(max_budget=0.06) + team: Final = scenario.team(models=[model]) + added: Final = candidate.request( + "POST", "/team/member_add", {"team_id": team, "member": {"user_id": user, "role": "user"}} + ) + assert added.status_code == 200, added.text + key: Final = scenario.key(team_id=team, user_id=user, models=[model]) + first: Final = candidate.chat(model, key=key, text=f"owner budget {uuid.uuid4().hex}") + assert first["usage"]["total_tokens"] == 40 + spent: Final = eventually( + lambda: read_rows('SELECT spend FROM "LiteLLM_UserTable" WHERE user_id=%s', (user,)), + lambda values: len(values) == 1 and float(values[0]["spend"]) >= 0.06, + seconds=70, + ) + assert float(spent[0]["spend"]) == pytest.approx(0.06) + assert read_rows( + 'SELECT max_budget FROM "LiteLLM_VerificationToken" WHERE token=%s', (sha256(key.encode()).hexdigest(),) + ) == [{"max_budget": None}] + upstream.get("/__observations").raise_for_status() + denied: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": f"over owner budget {uuid.uuid4().hex}"}]}, + key=key, + ) + assert denied.status_code == 422 and denied.json()["error"]["type"] == "budget_exceeded", denied.text + assert upstream.get("/__observations").json()["requests"] == [] From b41e6c966afb548a337ab0fe02feee1a7847bda3 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 23 Sep 2026 16:52:34 +0000 Subject: [PATCH 10/10] feat(models): add openrouter/stealth/space-bunny-alpha (#42759) Co-authored-by: kerry Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../model_prices_and_context_window_backup.json | 15 +++++++++++++++ model_prices_and_context_window.json | 15 +++++++++++++++ 2 files changed, 30 insertions(+) diff --git a/litellm/model_prices_and_context_window_backup.json b/litellm/model_prices_and_context_window_backup.json index cc3bf192908..da8013ae709 100644 --- a/litellm/model_prices_and_context_window_backup.json +++ b/litellm/model_prices_and_context_window_backup.json @@ -72194,6 +72194,21 @@ "supports_vision": false, "supports_web_search": false }, + "openrouter/stealth/space-bunny-alpha": { + "input_cost_per_token": 0, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 0, + "source": "https://openrouter.ai/stealth/space-bunny-alpha", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true + }, "openrouter/stepfun/step-3.5-flash": { "input_cost_per_token": 1e-07, "litellm_provider": "openrouter", diff --git a/model_prices_and_context_window.json b/model_prices_and_context_window.json index cc3bf192908..da8013ae709 100644 --- a/model_prices_and_context_window.json +++ b/model_prices_and_context_window.json @@ -72194,6 +72194,21 @@ "supports_vision": false, "supports_web_search": false }, + "openrouter/stealth/space-bunny-alpha": { + "input_cost_per_token": 0, + "litellm_provider": "openrouter", + "max_input_tokens": 1000000, + "max_output_tokens": 128000, + "max_tokens": 128000, + "mode": "chat", + "output_cost_per_token": 0, + "source": "https://openrouter.ai/stealth/space-bunny-alpha", + "supports_function_calling": true, + "supports_reasoning": true, + "supports_tool_choice": true, + "supports_video_input": true, + "supports_vision": true + }, "openrouter/stepfun/step-3.5-flash": { "input_cost_per_token": 1e-07, "litellm_provider": "openrouter",