diff --git a/litellm/integrations/asqav/asqav.py b/litellm/integrations/asqav/asqav.py index 5f0235fecd8..c287fef38f2 100644 --- a/litellm/integrations/asqav/asqav.py +++ b/litellm/integrations/asqav/asqav.py @@ -146,9 +146,9 @@ def _extract_loggable( if hasattr(response_obj, "choices") and response_obj.choices: finish_reason = response_obj.choices[0].finish_reason if hasattr(response_obj, "_hidden_params"): - provider_request_id = response_obj._hidden_params.get( - "x-request-id" - ) or response_obj._hidden_params.get("cf-ray") + provider_request_id = response_obj._hidden_params.get("x-request-id") or response_obj._hidden_params.get( + "cf-ray" + ) except Exception: # noqa: BLE001 pass @@ -172,9 +172,7 @@ def _extract_loggable( except Exception: # noqa: BLE001 pass if not call_id: - call_id = kwargs.get("litellm_call_id") or kwargs.get( - "id", str(int(time.time() * 1e6)) - ) + call_id = kwargs.get("litellm_call_id") or kwargs.get("id", str(int(time.time() * 1e6))) return { "call_id": call_id, @@ -220,13 +218,9 @@ class AsqavLogger(CustomLogger): ) -> None: super().__init__() - self._log_path: str = log_path or os.environ.get( - "ASQAV_LOG_PATH", _DEFAULT_LOG_PATH - ) + self._log_path: str = log_path or os.environ.get("ASQAV_LOG_PATH", _DEFAULT_LOG_PATH) self._redact_content: bool = ( - os.environ.get("ASQAV_REDACT_CONTENT", "true").lower() != "false" - if log_path is None - else redact_content + os.environ.get("ASQAV_REDACT_CONTENT", "true").lower() != "false" if log_path is None else redact_content ) self._lock: threading.Lock = threading.Lock() @@ -238,10 +232,7 @@ class AsqavLogger(CustomLogger): self._load_chain_tail() def __repr__(self) -> str: - return ( - f"AsqavLogger(log_path={self._log_path!r}," - f" redact_content={self._redact_content})" - ) + return f"AsqavLogger(log_path={self._log_path!r}, redact_content={self._redact_content})" # ------------------------------------------------------------------ # Chain state persistence @@ -265,9 +256,7 @@ class AsqavLogger(CustomLogger): self._prev_hash = last_record.get("record_hash", _GENESIS_HASH) self._call_count = last_record.get("seq", -1) + 1 except Exception: # noqa: BLE001 - verbose_logger.debug( - f"[AsqavLogger] Could not load chain tail: {traceback.format_exc()}" - ) + verbose_logger.debug(f"[AsqavLogger] Could not load chain tail: {traceback.format_exc()}") # ------------------------------------------------------------------ # Core record append @@ -287,18 +276,14 @@ class AsqavLogger(CustomLogger): """ start_time, end_time = timing try: - loggable = _extract_loggable( - kwargs, response_obj, start_time, end_time, status - ) + loggable = _extract_loggable(kwargs, response_obj, start_time, end_time, status) if not self._redact_content: # Store content in the clear when the operator explicitly opts in. loggable["messages"] = kwargs.get("messages") try: if hasattr(response_obj, "choices") and response_obj.choices: - loggable["response_content"] = response_obj.choices[ - 0 - ].message.content + loggable["response_content"] = response_obj.choices[0].message.content except Exception: # noqa: BLE001 pass @@ -327,9 +312,7 @@ class AsqavLogger(CustomLogger): self._call_count += 1 except Exception: # noqa: BLE001 - verbose_logger.debug( - f"[AsqavLogger] Unhandled error in _build_and_append: {traceback.format_exc()}" - ) + verbose_logger.debug(f"[AsqavLogger] Unhandled error in _build_and_append: {traceback.format_exc()}") def _write_record(self, record: dict[str, Any]) -> bool: """Append one record to the log file. Returns False if the write failed.""" @@ -360,9 +343,7 @@ class AsqavLogger(CustomLogger): raise return True except Exception: # noqa: BLE001 - verbose_logger.warning( - f"[AsqavLogger] Failed to write audit record: {traceback.format_exc()}" - ) + verbose_logger.warning(f"[AsqavLogger] Failed to write audit record: {traceback.format_exc()}") return False # ------------------------------------------------------------------ @@ -449,18 +430,14 @@ class AsqavLogger(CustomLogger): if computed_hash != stored_hash: return ( False, - f"line {lineno}: hash mismatch" - f" (stored={stored_hash[:12]}," - f" computed={computed_hash[:12]})", + f"line {lineno}: hash mismatch (stored={stored_hash[:12]}, computed={computed_hash[:12]})", ) rec_prev = record.get("prev_hash", _GENESIS_HASH) if rec_prev != prev_hash: return ( False, - f"line {lineno}: prev_hash chain break" - f" (expected={prev_hash[:12]}," - f" got={rec_prev[:12]})", + f"line {lineno}: prev_hash chain break (expected={prev_hash[:12]}, got={rec_prev[:12]})", ) prev_hash = stored_hash diff --git a/litellm/integrations/prometheus.py b/litellm/integrations/prometheus.py index 19c64aaa535..94c5ac525ce 100644 --- a/litellm/integrations/prometheus.py +++ b/litellm/integrations/prometheus.py @@ -2746,9 +2746,7 @@ class PrometheusLogger(CustomLogger): Set the deployment state. """ _labels = prometheus_label_factory( - supported_enum_labels=self.get_labels_for_metric( - metric_name="litellm_deployment_state" - ), + supported_enum_labels=self.get_labels_for_metric(metric_name="litellm_deployment_state"), enum_values=enum_values, ) self.litellm_deployment_state.labels(**_labels).set(state) @@ -2819,12 +2817,8 @@ class PrometheusLogger(CustomLogger): increment metric when litellm.Router / load balancing logic places a deployment in cool down """ _labels = prometheus_label_factory( - supported_enum_labels=self.get_labels_for_metric( - metric_name="litellm_deployment_cooled_down" - ), - enum_values=dataclasses.replace( - enum_values, exception_status=exception_status - ), + supported_enum_labels=self.get_labels_for_metric(metric_name="litellm_deployment_cooled_down"), + enum_values=dataclasses.replace(enum_values, exception_status=exception_status), ) self.litellm_deployment_cooled_down.labels(**_labels).inc() diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py index 598350701c3..6a332240bf2 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/streaming_iterator.py @@ -430,9 +430,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): will_merge_into_held = ( self.holding_stop_reason_chunk is not None and getattr(chunk, "usage", None) is not None ) - is_final_chunk = ( - bool(chunk.choices) and chunk.choices[0].finish_reason is not None - ) + is_final_chunk = bool(chunk.choices) and chunk.choices[0].finish_reason is not None processed_chunk = LiteLLMAnthropicMessagesAdapter().translate_streaming_openai_response_to_anthropic( response=chunk, current_content_block_index=self.current_content_block_index, @@ -655,9 +653,7 @@ class AnthropicStreamWrapper(AdapterCompletionStreamWrapper): will_merge_into_held = ( self.holding_stop_reason_chunk is not None and getattr(chunk, "usage", None) is not None ) - is_final_chunk = ( - bool(chunk.choices) and chunk.choices[0].finish_reason is not None - ) + is_final_chunk = bool(chunk.choices) and chunk.choices[0].finish_reason is not None processed_chunk = LiteLLMAnthropicMessagesAdapter().translate_streaming_openai_response_to_anthropic( response=chunk, current_content_block_index=self.current_content_block_index, diff --git a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py index bba0977b339..c858ca9b09e 100644 --- a/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py +++ b/litellm/llms/anthropic/experimental_pass_through/adapters/transformation.py @@ -1489,15 +1489,11 @@ class LiteLLMAnthropicMessagesAdapter: ## base case - final chunk w/ finish reason, or a usage-only chunk ## (choices=[]) that carries trailing usage. See #30761. - has_finish_reason = ( - bool(response.choices) and response.choices[0].finish_reason is not None - ) + has_finish_reason = bool(response.choices) and response.choices[0].finish_reason is not None has_usage_only_chunk = not response.choices and litellm_usage_chunk is not None if has_finish_reason or has_usage_only_chunk: stop_reason = ( - self._translate_openai_finish_reason_to_anthropic( - response.choices[0].finish_reason - ) + self._translate_openai_finish_reason_to_anthropic(response.choices[0].finish_reason) if response.choices else None ) diff --git a/litellm/llms/dashscope/cost_calculator.py b/litellm/llms/dashscope/cost_calculator.py index 45bd092c6c9..0182bf791b8 100644 --- a/litellm/llms/dashscope/cost_calculator.py +++ b/litellm/llms/dashscope/cost_calculator.py @@ -42,9 +42,7 @@ def _extract_token_breakdown(usage: Usage) -> TokenBreakdown: return TokenBreakdown(text_tokens, cached_tokens, completion_tokens, reasoning_tokens) -def _resolve_tier_cost_per_token( - tier: dict, cost_key: str, fallback_cost_key: str | None -) -> float: +def _resolve_tier_cost_per_token(tier: dict, cost_key: str, fallback_cost_key: str | None) -> float: """Resolve a tier's per-token cost. An explicit 0.0 is a real price (e.g. a free-cache-read tier) and must not @@ -118,9 +116,7 @@ def _calculate_tiered_cost( if tier_end > tier_start: tokens_in_tier = tier_end - tier_start - cost_per_token = _resolve_tier_cost_per_token( - tier, cost_key, fallback_cost_key - ) + cost_per_token = _resolve_tier_cost_per_token(tier, cost_key, fallback_cost_key) total_cost += tokens_in_tier * cost_per_token tokens_processed = tier_end @@ -129,9 +125,7 @@ def _calculate_tiered_cost( if tokens_processed < tokens and sorted_tiers: last_tier = sorted_tiers[-1] remaining_tokens = tokens - tokens_processed - cost_per_token = _resolve_tier_cost_per_token( - last_tier, cost_key, fallback_cost_key - ) + cost_per_token = _resolve_tier_cost_per_token(last_tier, cost_key, fallback_cost_key) total_cost += remaining_tokens * cost_per_token return total_cost diff --git a/litellm/llms/perplexity/search/transformation.py b/litellm/llms/perplexity/search/transformation.py index 1b82ed48805..198e22b6ace 100644 --- a/litellm/llms/perplexity/search/transformation.py +++ b/litellm/llms/perplexity/search/transformation.py @@ -118,11 +118,7 @@ class PerplexitySearchConfig(BaseSearchConfig): """ request_data: PerplexitySearchRequest = {"query": query} - forwarded = { - key: value - for key, value in optional_params.items() - if value is not None and key != "query" - } + forwarded = {key: value for key, value in optional_params.items() if value is not None and key != "query"} return dict(request_data, **forwarded) diff --git a/litellm/proxy/common_request_processing.py b/litellm/proxy/common_request_processing.py index 1df0afbded2..ab16c0e9acc 100644 --- a/litellm/proxy/common_request_processing.py +++ b/litellm/proxy/common_request_processing.py @@ -2454,11 +2454,7 @@ class ProxyBaseLLMRequestProcessing: """ logging_obj = request_data.get("litellm_logging_obj") response_chunks = getattr(response, "chunks", None) - chunks = ( - response_chunks - if isinstance(response_chunks, list) and response_chunks - else streamed_chunks - ) + chunks = response_chunks if isinstance(response_chunks, list) and response_chunks else streamed_chunks if logging_obj is None or not chunks: return first_chunk = chunks[0] @@ -2481,9 +2477,7 @@ class ProxyBaseLLMRequestProcessing: logging_obj=logging_obj, ) except Exception: # noqa: BLE001 - verbose_proxy_logger.exception( - "Failed to assemble partial streaming usage on client disconnect" - ) + verbose_proxy_logger.exception("Failed to assemble partial streaming usage on client disconnect") return if partial_response is None: return @@ -2510,9 +2504,7 @@ class ProxyBaseLLMRequestProcessing: prefer_async_handlers=True, ) except Exception: # noqa: BLE001 - verbose_proxy_logger.exception( - "Failed to record partial streaming usage on client disconnect" - ) + verbose_proxy_logger.exception("Failed to record partial streaming usage on client disconnect") @staticmethod async def async_streaming_data_generator( diff --git a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py index 2ce63cee89b..43691012aa1 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bedrock_guardrails.py @@ -82,9 +82,7 @@ _BEDROCK_INVOKE_GUARDRAIL_CHECKS_PATH = "/guardrail-checks/invoke" # more text blocks is split across multiple messages so ALL content is scanned -- # never truncated (truncation would let a user hide content past the limit). _BEDROCK_CHECKS_MAX_CONTENT_BLOCKS = 10 -_BEDROCK_CHECKS_KNOWN_KEYS = frozenset( - {"contentFilter", "promptAttack", "sensitiveInformation"} -) +_BEDROCK_CHECKS_KNOWN_KEYS = frozenset({"contentFilter", "promptAttack", "sensitiveInformation"}) # InvokeGuardrailChecks only accepts roles user/assistant/system. Map every other # OpenAI role onto one of these so NO message content is skipped (skipping would let # a user hide prohibited text in e.g. a tool/function message that the model still @@ -237,8 +235,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): # InvokeGuardrailChecks is detect-only: it never returns rewritten content, # so masking has no effect in checks mode. if self.checks is not None and ( - getattr(self, "mask_request_content", False) - or getattr(self, "mask_response_content", False) + getattr(self, "mask_request_content", False) or getattr(self, "mask_response_content", False) ): verbose_proxy_logger.warning( "Bedrock Guardrail: mask_request_content/mask_response_content have no " @@ -274,15 +271,11 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): sorted(_BEDROCK_CHECKS_KNOWN_KEYS), ) cleaned = { - key: value - for key, value in checks.items() - if key in _BEDROCK_CHECKS_KNOWN_KEYS and value is not None + key: value for key, value in checks.items() if key in _BEDROCK_CHECKS_KNOWN_KEYS and value is not None } return cleaned or None - def _create_bedrock_input_content_request( - self, messages: Optional[List[AllMessageValues]] - ) -> BedrockRequest: + def _create_bedrock_input_content_request(self, messages: Optional[List[AllMessageValues]]) -> BedrockRequest: """ Create a bedrock request for the input content - the LLM request. """ @@ -685,10 +678,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): # request_path for the resource-less InvokeGuardrailChecks endpoint (where # guardrailIdentifier/guardrailVersion are None and must not be interpolated). if request_path is None: - request_path = ( - f"/guardrail/{self.guardrailIdentifier}" - f"/version/{self.guardrailVersion}/apply" - ) + request_path = f"/guardrail/{self.guardrailIdentifier}/version/{self.guardrailVersion}/apply" proxy_endpoint_url = f"{proxy_endpoint_url}{request_path}" encoded_data = json.dumps(data).encode("utf-8") @@ -904,20 +894,14 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): guardrail_status="guardrail_failed_to_respond", start_time=start_time.timestamp(), end_time=datetime.now(timezone.utc).timestamp(), - duration=( - datetime.now(timezone.utc) - start_time - ).total_seconds(), + duration=(datetime.now(timezone.utc) - start_time).total_seconds(), event_type=event_type, ) - raise HTTPException( - status_code=status_code, detail=detail_message - ) from e + raise HTTPException(status_code=status_code, detail=detail_message) from e except HTTPException: raise # Endpoint down, timeout, or other HTTP/network errors - verbose_proxy_logger.error( - "Bedrock AI: failed to make guardrail request: %s", str(e) - ) + verbose_proxy_logger.error("Bedrock AI: failed to make guardrail request: %s", str(e)) self.add_standard_logging_guardrail_information_to_request_data( guardrail_provider=self.guardrail_provider, guardrail_json_response={"error": str(e)}, @@ -933,9 +917,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ########### InvokeGuardrailChecks (resource-less, detect-only) ############ @staticmethod - def _chunk_texts_into_checks_messages( - role: str, texts: list[str] - ) -> list[BedrockChecksMessage]: + def _chunk_texts_into_checks_messages(role: str, texts: list[str]) -> list[BedrockChecksMessage]: """Group ``texts`` into role-tagged messages of <= the API content-block cap. A source message with more text blocks than the per-message limit is split @@ -973,9 +955,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): # Reuse the ApplyGuardrail output extractor (single source of truth for # pulling assistant text out of a ModelResponse), then re-tag as an # assistant turn for the role-based InvokeGuardrailChecks payload. - output_request = self._create_bedrock_output_content_request( - response=response - ) + output_request = self._create_bedrock_output_content_request(response=response) output_texts: list[str] = [] for item in output_request.get("content") or []: text = (item.get("text") or {}).get("text") @@ -1028,14 +1008,10 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): api_key=api_key, request_path=_BEDROCK_INVOKE_GUARDRAIL_CHECKS_PATH, ) - verbose_proxy_logger.debug( - "Bedrock InvokeGuardrailChecks request url: %s", prepared_request.url - ) + verbose_proxy_logger.debug("Bedrock InvokeGuardrailChecks request url: %s", prepared_request.url) event_type = logging_event_type or ( - GuardrailEventHooks.pre_call - if source == "INPUT" - else GuardrailEventHooks.post_call + GuardrailEventHooks.pre_call if source == "INPUT" else GuardrailEventHooks.post_call ) httpx_response = await self._sign_and_post( @@ -1046,9 +1022,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ) if httpx_response.status_code != 200: - status_code, detail_message = self._parse_bedrock_guardrail_error_response( - httpx_response - ) + status_code, detail_message = self._parse_bedrock_guardrail_error_response(httpx_response) verbose_proxy_logger.error( "Bedrock InvokeGuardrailChecks: error response. Status %s: %s", httpx_response.status_code, @@ -1066,18 +1040,14 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ) raise HTTPException(status_code=status_code, detail=detail_message) - json_response: BedrockGuardrailChecksResponse = cast( - BedrockGuardrailChecksResponse, httpx_response.json() - ) + json_response: BedrockGuardrailChecksResponse = cast(BedrockGuardrailChecksResponse, httpx_response.json()) violations = self._collect_invoke_checks_violations(json_response) # Log a copy with PII location offsets stripped: offsets + the (separately # logged) request messages would otherwise reconstruct the detected PII span. self.add_standard_logging_guardrail_information_to_request_data( guardrail_provider=self.guardrail_provider, - guardrail_json_response=self._sanitize_invoke_checks_response_for_logging( - json_response - ), + guardrail_json_response=self._sanitize_invoke_checks_response_for_logging(json_response), request_data=request_data or {}, guardrail_status=self._get_invoke_checks_status(bool(violations)), start_time=start_time.timestamp(), @@ -1167,9 +1137,7 @@ class BedrockGuardrail(CustomGuardrail, BaseAWSLLM): ] if categories: tracing_detail["violation_categories"] = categories - tracing_detail["guardrail_action"] = ( - "GUARDRAIL_INTERVENED" if violations else "NONE" - ) + tracing_detail["guardrail_action"] = "GUARDRAIL_INTERVENED" if violations else "NONE" return tracing_detail def _get_block_exception_for_checks( diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/__init__.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/__init__.py index 4bdb12566e8..298ed6b7b56 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/__init__.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/__init__.py @@ -29,17 +29,13 @@ def initialize_guardrail( config=GuardrailConfig( thresholds=RiskThresholds( bias_threshold=getattr(litellm_params, "bias_threshold", 0.5), - hallucination_threshold=getattr( - litellm_params, "hallucination_threshold", 0.5 - ), + hallucination_threshold=getattr(litellm_params, "hallucination_threshold", 0.5), flag_threshold=getattr(litellm_params, "risk_flag_threshold", 0.25), block_threshold=getattr(litellm_params, "risk_block_threshold", 0.5), ), weights=RiskWeights( bias_weight=getattr(litellm_params, "bias_weight", 0.4), - hallucination_weight=getattr( - litellm_params, "hallucination_weight", 0.6 - ), + hallucination_weight=getattr(litellm_params, "hallucination_weight", 0.6), ), behavior=GuardrailBehaviorConfig( block_on_high_risk=getattr(litellm_params, "block_on_high_risk", True), @@ -49,15 +45,11 @@ def initialize_guardrail( violation_message=getattr(litellm_params, "violation_message", None), ), session=GuardrailSessionConfig( - violation_message_template=getattr( - litellm_params, "violation_message_template", None - ), + violation_message_template=getattr(litellm_params, "violation_message_template", None), ), ), ) - logging_callback_manager.add_litellm_callback( - callback - ) # pyright: ignore[reportUnknownMemberType] + logging_callback_manager.add_litellm_callback(callback) # pyright: ignore[reportUnknownMemberType] return callback diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/bias_hallucination_estimator.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/bias_hallucination_estimator.py index 1cfd50f47c8..c270ff95b1e 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/bias_hallucination_estimator.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/bias_hallucination_estimator.py @@ -48,9 +48,7 @@ GuardrailEventHookInput = Union[ Sequence[str], Mode, ] -NormalizedGuardrailEventHook = Union[ - GuardrailEventHooks, list[GuardrailEventHooks], Mode -] +NormalizedGuardrailEventHook = Union[GuardrailEventHooks, list[GuardrailEventHooks], Mode] class FunctionLike(Protocol): @@ -102,12 +100,8 @@ class GuardrailConfig: thresholds: RiskThresholds = dataclasses_field(default_factory=RiskThresholds) weights: RiskWeights = dataclasses_field(default_factory=RiskWeights) - behavior: GuardrailBehaviorConfig = dataclasses_field( - default_factory=GuardrailBehaviorConfig - ) - session: GuardrailSessionConfig = dataclasses_field( - default_factory=GuardrailSessionConfig - ) + behavior: GuardrailBehaviorConfig = dataclasses_field(default_factory=GuardrailBehaviorConfig) + session: GuardrailSessionConfig = dataclasses_field(default_factory=GuardrailSessionConfig) class BiasHallucinationEstimatorGuardrail(CustomGuardrail): @@ -130,8 +124,7 @@ class BiasHallucinationEstimatorGuardrail(CustomGuardrail): GuardrailEventHooks.pre_call, GuardrailEventHooks.post_call, ], - event_hook=self._normalize_event_hook(event_hook) - or GuardrailEventHooks.post_call, + event_hook=self._normalize_event_hook(event_hook) or GuardrailEventHooks.post_call, mask_request_content=_session.mask_request_content, mask_response_content=_session.mask_response_content, violation_message_template=_session.violation_message_template, @@ -174,13 +167,9 @@ class BiasHallucinationEstimatorGuardrail(CustomGuardrail): return inputs analyses = tuple(self._analyze_text(text=text) for text in texts) - highest_risk = max( - analyses, key=lambda analysis: analysis.risk.overall_risk_percentage - ) + highest_risk = max(analyses, key=lambda analysis: analysis.risk.overall_risk_percentage) decision = self._decision(highest_risk.risk.recommendation) - status: GuardrailStatus = ( - "guardrail_intervened" if decision == "blocked" else "success" - ) + status: GuardrailStatus = "guardrail_intervened" if decision == "blocked" else "success" response_payload = self._build_response_payload( analyses=analyses, decision=decision, @@ -247,15 +236,9 @@ class BiasHallucinationEstimatorGuardrail(CustomGuardrail): return { "decision": decision, "input_type": input_type, - "risk_scores": [ - analysis.risk.model_dump(mode="json") for analysis in analyses - ], + "risk_scores": [analysis.risk.model_dump(mode="json") for analysis in analyses], "bias": [ - { - k: v - for k, v in analysis.bias.model_dump(mode="json").items() - if k not in _LOG_EXCLUDED_FIELDS - } + {k: v for k, v in analysis.bias.model_dump(mode="json").items() if k not in _LOG_EXCLUDED_FIELDS} for analysis in analyses ], "hallucination": [ @@ -306,9 +289,7 @@ class BiasHallucinationEstimatorGuardrail(CustomGuardrail): @staticmethod def _detection_methods(response_payload: dict[str, object]) -> str | None: - categories = BiasHallucinationEstimatorGuardrail._violation_categories( - response_payload - ) + categories = BiasHallucinationEstimatorGuardrail._violation_categories(response_payload) if not categories: return None return "regex,keyword" @@ -323,9 +304,7 @@ class BiasHallucinationEstimatorGuardrail(CustomGuardrail): dict.fromkeys( issue.split(":", 1)[0] for risk_score in typed_risk_scores - for issue in BiasHallucinationEstimatorGuardrail._detected_issues( - risk_score - ) + for issue in BiasHallucinationEstimatorGuardrail._detected_issues(risk_score) ) ) @@ -370,9 +349,7 @@ class BiasHallucinationEstimatorGuardrail(CustomGuardrail): tool_call_texts = tuple( text for tool_call in inputs.get("tool_calls", []) - for text in ( - BiasHallucinationEstimatorGuardrail._tool_call_text(tool_call), - ) + for text in (BiasHallucinationEstimatorGuardrail._tool_call_text(tool_call),) if text ) return texts + tool_call_texts diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/data_sources.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/data_sources.py index 5ed7530cbae..a746d10ba65 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/data_sources.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/data_sources.py @@ -52,9 +52,7 @@ class DataSourceResult: self.metadata = metadata or {} def __repr__(self) -> str: - return ( - f"DataSourceResult(source={self.source}, confidence={self.confidence:.2f})" - ) + return f"DataSourceResult(source={self.source}, confidence={self.confidence:.2f})" class DataSource(ABC): @@ -64,9 +62,7 @@ class DataSource(ABC): self.priority = priority @abstractmethod - async def search( - self, query: str, limit: int = 5 - ) -> list[DataSourceResult]: ... # pragma: no cover + async def search(self, query: str, limit: int = 5) -> list[DataSourceResult]: ... # pragma: no cover async def verify_fact(self, claim: str) -> tuple[bool, str | None]: results = await self.search(claim, limit=1) @@ -112,9 +108,7 @@ def _keyword_search( text = _get_doc_text(doc) matching = len(set(re.findall(r"\b\w+\b", text.lower())) & query_words) score = matching / len(query_words) - scored.append( - (score, DataSourceResult(text=text, source=source_name, confidence=score)) - ) + scored.append((score, DataSourceResult(text=text, source=source_name, confidence=score))) scored.sort(reverse=True, key=lambda x: x[0]) return [result for _, result in scored[:limit]] @@ -150,9 +144,7 @@ class FileDataSource(DataSource): return [] if path.suffix == ".json": with open(path) as f: - return cast( - list[str | dict[str, Any]], _parse_json_docs(json.load(f)) - ) + return cast(list[str | dict[str, Any]], _parse_json_docs(json.load(f))) if path.suffix in {".csv", ".txt"}: with open(path) as f: return cast( @@ -197,9 +189,7 @@ class URLDataSource(DataSource): async def _fetch_all(self) -> list[str | dict[str, Any]]: if not _AIOHTTP_AVAILABLE: return [] - results = await asyncio.gather( - *[self._fetch_url(u) for u in self.urls], return_exceptions=True - ) + results = await asyncio.gather(*[self._fetch_url(u) for u in self.urls], return_exceptions=True) return [doc for batch in results if isinstance(batch, list) for doc in batch] async def _fetch_url(self, url: str) -> list[str | dict[str, Any]]: @@ -207,9 +197,7 @@ class URLDataSource(DataSource): import aiohttp # type: ignore[import-untyped] async with aiohttp.ClientSession() as session: - async with session.get( - url, timeout=aiohttp.ClientTimeout(total=self.timeout) - ) as response: + async with session.get(url, timeout=aiohttp.ClientTimeout(total=self.timeout)) as response: if response.status == 200: content = await response.text() return self._parse_content(content) @@ -220,9 +208,7 @@ class URLDataSource(DataSource): @staticmethod def _parse_content(content: str) -> list[str | dict[str, Any]]: try: - return cast( - list[str | dict[str, Any]], _parse_json_docs(json.loads(content)) - ) + return cast(list[str | dict[str, Any]], _parse_json_docs(json.loads(content))) except json.JSONDecodeError: return cast(list[str | dict[str, Any]], [{"text": content}]) @@ -236,9 +222,7 @@ class ContextDocumentDataSource(DataSource): priority: int = 100, ) -> None: super().__init__(name=name, enabled=enabled, priority=priority) - self._documents: list[str | dict[str, Any]] = cast( - list[str | dict[str, Any]], documents - ) + self._documents: list[str | dict[str, Any]] = cast(list[str | dict[str, Any]], documents) self._index = _build_keyword_index(self._documents) async def search(self, query: str, limit: int = 5) -> list[DataSourceResult]: @@ -314,9 +298,7 @@ class VectorStoreDataSource(DataSource): if not embedding: return [] if self.provider == "pinecone": - results = self.client.query( - embedding, top_k=limit, include_metadata=True - ) + results = self.client.query(embedding, top_k=limit, include_metadata=True) return [ DataSourceResult( text=match.get("metadata", {}).get("text", ""), @@ -333,12 +315,7 @@ class VectorStoreDataSource(DataSource): .do() ) docs = response.get("data", {}).get("Get", {}) - return [ - DataSourceResult( - text=doc.get("text", ""), source=self.name, confidence=0.8 - ) - for doc in docs - ] + return [DataSourceResult(text=doc.get("text", ""), source=self.name, confidence=0.8) for doc in docs] except Exception: # noqa: BLE001 pass return [] @@ -386,12 +363,8 @@ class KnowledgeGraphDataSource(DataSource): if response.status == 200: data = await response.json() return [ - DataSourceResult( - text=text, source=self.name, confidence=0.9 - ) - for binding in data.get("results", {}).get("bindings", [])[ - :limit - ] + DataSourceResult(text=text, source=self.name, confidence=0.9) + for binding in data.get("results", {}).get("bindings", [])[:limit] for text in (self._extract_text(binding),) if text ] @@ -437,7 +410,5 @@ class FactCheckDataSource(DataSource): self.api_key = api_key self.config = config - async def search( - self, query: str, limit: int = 5 - ) -> list[DataSourceResult]: # noqa: ARG002 + async def search(self, query: str, limit: int = 5) -> list[DataSourceResult]: # noqa: ARG002 return [] diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/estimator_core.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/estimator_core.py index 86074d6249a..2e11f99ac94 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/estimator_core.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/estimator_core.py @@ -35,9 +35,7 @@ class BiasDetector: ) @staticmethod - def _match_patterns( - text: str, patterns: tuple[PatternRule, ...] - ) -> tuple[tuple[str, str, float], ...]: + def _match_patterns(text: str, patterns: tuple[PatternRule, ...]) -> tuple[tuple[str, str, float], ...]: return tuple( (rule.name, clip_example(match.group(0)), rule.score) for rule in patterns @@ -59,9 +57,7 @@ class HallucinationDetector: sentences = split_sentences(text) unsourced_claims = self._find_unsourced_statistics(sentences) missing_citations = self._find_rule_matches(text, CITATION_GAP_PATTERNS) - fabricated_specificity = self._find_rule_matches( - text, FABRICATED_SPECIFICITY_PATTERNS - ) + fabricated_specificity = self._find_rule_matches(text, FABRICATED_SPECIFICITY_PATTERNS) patterns_found = self._patterns_found( has_unsourced_claims=bool(unsourced_claims), has_missing_citations=bool(missing_citations), @@ -73,9 +69,7 @@ class HallucinationDetector: fabricated_specificity=fabricated_specificity, ) examples = unique_preserve_order( - tuple(unsourced_claims) - + tuple(missing_citations) - + tuple(fabricated_specificity) + tuple(unsourced_claims) + tuple(missing_citations) + tuple(fabricated_specificity) ) return HallucinationAnalysis( @@ -99,13 +93,9 @@ class HallucinationDetector: ) @staticmethod - def _find_rule_matches( - text: str, rules: tuple[PatternRule, ...] - ) -> tuple[str, ...]: + def _find_rule_matches(text: str, rules: tuple[PatternRule, ...]) -> tuple[str, ...]: return unique_preserve_order( - clip_example(match.group(0)) - for rule in rules - for match in rule.pattern.finditer(text) + clip_example(match.group(0)) for rule in rules for match in rule.pattern.finditer(text) ) @staticmethod diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/grounding_checker.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/grounding_checker.py index 3fab770ee06..7220fc76c0c 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/grounding_checker.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/grounding_checker.py @@ -38,15 +38,11 @@ class GroundingChecker: claim_elements = self._extract_verifiable_elements(claim) if not any(claim_elements.values()): - return GroundingResult( - claim=claim, reasoning="No verifiable elements found in claim" - ) + return GroundingResult(claim=claim, reasoning="No verifiable elements found in claim") enabled_sources = [ds for ds in self.data_sources if ds.enabled] if not enabled_sources: - return GroundingResult( - claim=claim, reasoning="No enabled data sources available" - ) + return GroundingResult(claim=claim, reasoning="No enabled data sources available") raw_results = await asyncio.gather( *[self._search_source_safe(source, claim) for source in enabled_sources], @@ -87,21 +83,13 @@ class GroundingChecker: ) async def verify_multiple_claims(self, claims: list[str]) -> list[GroundingResult]: - return list( - await asyncio.gather(*[self.check_claim_grounding(c) for c in claims]) - ) + return list(await asyncio.gather(*[self.check_claim_grounding(c) for c in claims])) - async def _search_source_safe( - self, source: DataSource, claim: str - ) -> list[DataSourceResult]: + async def _search_source_safe(self, source: DataSource, claim: str) -> list[DataSourceResult]: try: - return await asyncio.wait_for( - source.search(claim, limit=3), timeout=self.timeout_per_source - ) + return await asyncio.wait_for(source.search(claim, limit=3), timeout=self.timeout_per_source) except asyncio.TimeoutError: - verbose_logger.warning( - f"Timeout searching {source.name} for claim: {claim}" - ) + verbose_logger.warning(f"Timeout searching {source.name} for claim: {claim}") return [] except Exception as e: # noqa: BLE001 verbose_logger.warning(f"Error searching {source.name}: {e}") @@ -135,11 +123,7 @@ class GroundingChecker: numbers = re.findall(r"\b\d+(?:,\d{3})*(?:\.\d+)?\s?%?|\b\d{4}\b", claim) dates = re.findall(r"\b(?:\d{1,2}/\d{1,2}/\d{4}|\d{4})\b", claim) entities = re.findall(r"\b[A-Z][a-z]+(?:\s+[A-Z][a-z]+)*\b", claim) - keywords = [ - w - for w in re.findall(r"\b[a-z]+\b", claim.lower()) - if w not in stop_words and len(w) > 3 - ] + keywords = [w for w in re.findall(r"\b[a-z]+\b", claim.lower()) if w not in stop_words and len(w) > 3] return { "numbers": numbers[:5], "dates": dates[:3], @@ -148,24 +132,16 @@ class GroundingChecker: } @staticmethod - def _boost_confidence( - result: DataSourceResult, claim_elements: dict[str, list[str]] - ) -> float: + def _boost_confidence(result: DataSourceResult, claim_elements: dict[str, list[str]]) -> float: boost = 0.0 - result_numbers = set( - re.findall(r"\b\d+(?:,\d{3})*(?:\.\d+)?\s?%?|\b\d{4}\b", result.text) - ) + result_numbers = set(re.findall(r"\b\d+(?:,\d{3})*(?:\.\d+)?\s?%?|\b\d{4}\b", result.text)) if set(claim_elements.get("numbers", [])) & result_numbers: boost += 0.1 - result_entities = set( - re.findall(r"\b[A-Z][a-z]+(?:\s+[A-Z][a-z]+)*\b", result.text) - ) + result_entities = set(re.findall(r"\b[A-Z][a-z]+(?:\s+[A-Z][a-z]+)*\b", result.text)) if set(claim_elements.get("entities", [])) & result_entities: boost += 0.15 result_words = set(re.findall(r"\b[a-z]+\b", result.text.lower())) claim_keywords = set(claim_elements.get("keywords", [])) if claim_keywords and result_words: - boost += min( - 0.2, len(claim_keywords & result_words) / len(claim_keywords) * 0.25 - ) + boost += min(0.2, len(claim_keywords & result_words) / len(claim_keywords) * 0.25) return round(min(1.0, result.confidence + boost), 2) diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/risk_scorer.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/risk_scorer.py index e0d7fdc5c5c..48ae7b00c32 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/risk_scorer.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/risk_scorer.py @@ -81,11 +81,7 @@ class RiskScorer: if total_weight <= 0: return 0.0 return min( - ( - bias_score * self.bias_weight - + hallucination_score * self.hallucination_weight - ) - / total_weight, + (bias_score * self.bias_weight + hallucination_score * self.hallucination_weight) / total_weight, 1.0, ) diff --git a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/utils.py b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/utils.py index 4d7ad5a8ebe..1f2f381a12f 100644 --- a/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/utils.py +++ b/litellm/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator/utils.py @@ -10,11 +10,7 @@ def split_sentences(text: str) -> tuple[str, ...]: normalized_text = " ".join(text.split()) if not normalized_text: return () - return tuple( - sentence.strip() - for sentence in SENTENCE_SPLIT_PATTERN.split(normalized_text) - if sentence.strip() - ) + return tuple(sentence.strip() for sentence in SENTENCE_SPLIT_PATTERN.split(normalized_text) if sentence.strip()) def clip_example(text: str, max_length: int = 160) -> str: diff --git a/litellm/proxy/guardrails/usage_endpoints.py b/litellm/proxy/guardrails/usage_endpoints.py index e2d26afa549..d93babce0b2 100644 --- a/litellm/proxy/guardrails/usage_endpoints.py +++ b/litellm/proxy/guardrails/usage_endpoints.py @@ -101,9 +101,7 @@ def _trend_from_comparison(current_fail: float, previous_fail: float) -> str: return "stable" -def _aggregate_daily_metrics( - metrics: list[Any], id_attr: str -) -> Dict[str, Dict[str, Any]]: +def _aggregate_daily_metrics(metrics: list[Any], id_attr: str) -> Dict[str, Dict[str, Any]]: agg: Dict[str, Dict[str, Any]] = {} for m in metrics: gid = getattr(m, id_attr) @@ -172,8 +170,7 @@ def _get_config_loaded_guardrails() -> list[Any]: return [ guardrail for guardrail in IN_MEMORY_GUARDRAIL_HANDLER.list_in_memory_guardrails() - if IN_MEMORY_GUARDRAIL_HANDLER.get_source(guardrail.get("guardrail_id") or "") - == "config" + if IN_MEMORY_GUARDRAIL_HANDLER.get_source(guardrail.get("guardrail_id") or "") == "config" ] @@ -324,9 +321,7 @@ async def guardrails_usage_overview( try: # Guardrails from DB - guardrails = _merge_config_loaded_guardrails( - await GuardrailsRepository(prisma_client).table.find_many() - ) + guardrails = _merge_config_loaded_guardrails(await GuardrailsRepository(prisma_client).table.find_many()) # Daily metrics in range metrics = await DailyGuardrailMetricsRepository(prisma_client).table.find_many( @@ -478,9 +473,7 @@ def _build_usage_logs_where( return where -def _usage_log_entry_from_row( - r: object, sl: object, action_filter: str | None -) -> Optional[UsageLogEntry]: +def _usage_log_entry_from_row(r: object, sl: object, action_filter: str | None) -> Optional[UsageLogEntry]: meta = sl.metadata if isinstance(meta, str): try: @@ -611,11 +604,7 @@ async def guardrails_usage_logs( ) or _find_config_loaded_guardrail(guardrail_id) if guardrail: logical_name = _get_guardrail_dict_field(guardrail, "guardrail_name") - if ( - logical_name - and isinstance(logical_name, str) - and logical_name not in effective_guardrail_ids - ): + if logical_name and isinstance(logical_name, str) and logical_name not in effective_guardrail_ids: effective_guardrail_ids.append(logical_name) where = _build_usage_logs_where(effective_guardrail_ids or None, policy_id, start_date, end_date) diff --git a/litellm/proxy/management_endpoints/ui_sso.py b/litellm/proxy/management_endpoints/ui_sso.py index 7777a0efe23..4a1cdb7d5fc 100644 --- a/litellm/proxy/management_endpoints/ui_sso.py +++ b/litellm/proxy/management_endpoints/ui_sso.py @@ -138,16 +138,12 @@ router = APIRouter() # response convertors see the same fields in the PKCE path as in the non-PKCE path. _OAUTH_TOKEN_FIELDS = frozenset({"access_token", "id_token", "refresh_token"}) _CLI_SSO_FLOW_CACHE_KEY_PREFIX = f"{CLI_SSO_SESSION_CACHE_KEY_PREFIX}:flow" -_CLI_SSO_START_RATE_LIMIT_CACHE_KEY_PREFIX = ( - f"{_CLI_SSO_FLOW_CACHE_KEY_PREFIX}:start_rate_limit" -) +_CLI_SSO_START_RATE_LIMIT_CACHE_KEY_PREFIX = f"{_CLI_SSO_FLOW_CACHE_KEY_PREFIX}:start_rate_limit" _CLI_SSO_START_RATE_LIMIT_WINDOW_SECONDS = 60 _CLI_SSO_START_RATE_LIMIT_MAX_ATTEMPTS = 30 _CLI_SSO_USER_CODE_ALPHABET = "ABCDEFGHJKLMNPQRSTUVWXYZ23456789" _CLI_SSO_LOGIN_ID_RE = re.compile(r"^cli-[A-Za-z0-9_-]{12,124}$") -_CLI_SSO_USER_CODE_RE = re.compile( - rf"^[{_CLI_SSO_USER_CODE_ALPHABET}]{{4}}-[{_CLI_SSO_USER_CODE_ALPHABET}]{{4}}$" -) +_CLI_SSO_USER_CODE_RE = re.compile(rf"^[{_CLI_SSO_USER_CODE_ALPHABET}]{{4}}-[{_CLI_SSO_USER_CODE_ALPHABET}]{{4}}$") _CLI_SSO_SCALAR_TYPES = (str, int, float, bool) _CLI_SSO_DEST_KEY_RE = re.compile(r"^[A-Za-z0-9_.-]+$") _CLI_SSO_SECRET_KEY_FRAGMENTS = frozenset( @@ -186,9 +182,7 @@ def _is_valid_cli_sso_login_id(login_id: Optional[str]) -> bool: def _is_valid_cli_sso_user_code(user_code: str | None) -> bool: - return isinstance(user_code, str) and bool( - _CLI_SSO_USER_CODE_RE.fullmatch(user_code) - ) + return isinstance(user_code, str) and bool(_CLI_SSO_USER_CODE_RE.fullmatch(user_code)) def _cli_sso_verification_uri_complete_enabled() -> bool: @@ -220,15 +214,8 @@ def _cli_sso_start_response_body( } -def _get_cli_sso_start_rate_limit_cache_key( - request: Request, use_x_forwarded_for: Optional[bool] = False -) -> str: - client_ip = ( - _get_request_ip_address( - request=request, use_x_forwarded_for=use_x_forwarded_for - ) - or "unknown" - ) +def _get_cli_sso_start_rate_limit_cache_key(request: Request, use_x_forwarded_for: Optional[bool] = False) -> str: + client_ip = _get_request_ip_address(request=request, use_x_forwarded_for=use_x_forwarded_for) or "unknown" client_ip_hash = _hash_cli_sso_secret(client_ip) return f"{_CLI_SSO_START_RATE_LIMIT_CACHE_KEY_PREFIX}:{client_ip_hash}" @@ -274,9 +261,7 @@ def _set_cli_sso_flow(login_id: str, cache: DualCache, flow: dict) -> None: def _verify_cli_sso_poll_secret(flow: dict, poll_secret: Optional[str]) -> bool: expected_poll_secret_hash = flow.get("poll_secret_hash") - if not isinstance(expected_poll_secret_hash, str) or not isinstance( - poll_secret, str - ): + if not isinstance(expected_poll_secret_hash, str) or not isinstance(poll_secret, str): return False supplied_poll_secret_hash = _hash_cli_sso_secret(poll_secret) return secrets.compare_digest(supplied_poll_secret_hash, expected_poll_secret_hash) @@ -361,9 +346,7 @@ def _get_nested_claim_value(data: Dict[str, Any], claim_path: str) -> Any: return current -def _extract_sso_claim_value( - result: Union[CustomOpenID, OpenID, dict], claim_path: str -) -> Any: +def _extract_sso_claim_value(result: Union[CustomOpenID, OpenID, dict], claim_path: str) -> Any: extra_fields = getattr(result, "extra_fields", None) if isinstance(extra_fields, dict): if claim_path in extra_fields: @@ -379,9 +362,7 @@ def _extract_sso_claim_value( return _get_nested_claim_value(result_dict, claim_path) -def _set_nested_metadata_value( - metadata: Dict[str, Any], key_path: str, value: Any -) -> None: +def _set_nested_metadata_value(metadata: Dict[str, Any], key_path: str, value: Any) -> None: placeholder = "\x00" parts = key_path.replace("\\.", placeholder).split(".") parts = [p.replace(placeholder, ".") for p in parts] @@ -428,18 +409,14 @@ def build_cli_sso_attribution_metadata( metadata: Dict[str, Any] = {} for source_claim, dest_key in claim_map: if not _is_safe_cli_sso_metadata_dest_key(dest_key): - verbose_proxy_logger.debug( - f"Skipping unsafe CLI SSO metadata destination key: {dest_key}" - ) + verbose_proxy_logger.debug(f"Skipping unsafe CLI SSO metadata destination key: {dest_key}") continue raw_value = _extract_sso_claim_value(result=result, claim_path=source_claim) if not _is_safe_cli_sso_scalar_claim_value(raw_value): continue - _set_nested_metadata_value( - metadata=metadata, key_path=dest_key, value=raw_value - ) + _set_nested_metadata_value(metadata=metadata, key_path=dest_key, value=raw_value) return metadata @@ -454,9 +431,7 @@ def _merge_cli_sso_attribution_metadata( are merged iteratively so attribution claims do not clobber unrelated keys under the same parent. """ - pending: List[Tuple[Dict[str, Any], Dict[str, Any]]] = [ - (existing_metadata, attribution_metadata) - ] + pending: List[Tuple[Dict[str, Any], Dict[str, Any]]] = [(existing_metadata, attribution_metadata)] while pending: target, source = pending.pop() for key, value in source.items(): @@ -479,9 +454,7 @@ async def _persist_cli_sso_user_metadata( return try: - user_row = await UserRepository(prisma_client).table.find_unique( - where={"user_id": user_id} - ) + user_row = await UserRepository(prisma_client).table.find_unique(where={"user_id": user_id}) existing_metadata: Dict[str, Any] = {} if user_row is not None: row_metadata = user_row.metadata @@ -501,9 +474,7 @@ async def _persist_cli_sso_user_metadata( f"{list(_flatten_cli_sso_metadata_for_poll(attribution_metadata).keys())}" ) except Exception as e: - verbose_proxy_logger.error( - f"Failed to persist CLI SSO attribution metadata for user {user_id}: {e}" - ) + verbose_proxy_logger.error(f"Failed to persist CLI SSO attribution metadata for user {user_id}: {e}") def _cli_poll_attribution_metadata_from_session( @@ -522,9 +493,7 @@ def _render_cli_sso_verification_page( ) -> str: escaped_verify_url = escape(verify_url, quote=True) escaped_browser_complete_token = escape(browser_complete_token, quote=True) - user_code_value_attr = ( - f' value="{escape(prefill_user_code, quote=True)}"' if prefill_user_code else "" - ) + user_code_value_attr = f' value="{escape(prefill_user_code, quote=True)}"' if prefill_user_code else "" instructions = ( "Confirm the verification code below to finish this login." if prefill_user_code @@ -603,9 +572,7 @@ async def cli_sso_start(request: Request): _check_cli_sso_start_rate_limit( request=request, cache=user_api_key_cache, - use_x_forwarded_for=bool( - (general_settings or {}).get("use_x_forwarded_for", False) - ), + use_x_forwarded_for=bool((general_settings or {}).get("use_x_forwarded_for", False)), ) login_id = f"cli-{secrets.token_urlsafe(24)}" @@ -623,9 +590,7 @@ async def cli_sso_start(request: Request): verification_uri_complete: str | None = ( ( - get_custom_url( - request_base_url=str(request.base_url), route="sso/key/generate" - ) + get_custom_url(request_base_url=str(request.base_url), route="sso/key/generate") + "?" + urlencode( { @@ -646,9 +611,7 @@ async def cli_sso_start(request: Request): ) -@router.post( - "/sso/cli/complete/{login_id}", tags=["experimental"], include_in_schema=False -) +@router.post("/sso/cli/complete/{login_id}", tags=["experimental"], include_in_schema=False) async def cli_sso_complete(request: Request, login_id: str): from fastapi.responses import HTMLResponse @@ -664,15 +627,9 @@ async def cli_sso_complete(request: Request, login_id: str): body = (await request.body()).decode("utf-8") form_values = parse_qs(body) supplied_user_code = (form_values.get("user_code") or [""])[0] - supplied_browser_complete_token = ( - form_values.get("browser_complete_token") or [""] - )[0] - supplied_user_code_hash = _hash_cli_sso_secret( - _normalize_cli_sso_user_code(supplied_user_code) - ) - supplied_browser_complete_token_hash = _hash_cli_sso_secret( - supplied_browser_complete_token - ) + supplied_browser_complete_token = (form_values.get("browser_complete_token") or [""])[0] + supplied_user_code_hash = _hash_cli_sso_secret(_normalize_cli_sso_user_code(supplied_user_code)) + supplied_browser_complete_token_hash = _hash_cli_sso_secret(supplied_browser_complete_token) expected_user_code_hash = flow.get("user_code_hash") if not isinstance(expected_user_code_hash, str) or not secrets.compare_digest( @@ -681,9 +638,7 @@ async def cli_sso_complete(request: Request, login_id: str): raise HTTPException(status_code=400, detail="Invalid verification code") expected_browser_complete_token_hash = flow.get("browser_complete_token_hash") - if not isinstance( - expected_browser_complete_token_hash, str - ) or not secrets.compare_digest( + if not isinstance(expected_browser_complete_token_hash, str) or not secrets.compare_digest( supplied_browser_complete_token_hash, expected_browser_complete_token_hash ): raise HTTPException(status_code=400, detail="Invalid verification code") @@ -754,9 +709,7 @@ def determine_role_from_groups( for role in role_hierarchy: if role in role_mappings.roles: role_groups = role_mappings.roles[role] - if isinstance(role_groups, list) and user_groups_set.intersection( - set(role_groups) - ): + if isinstance(role_groups, list) and user_groups_set.intersection(set(role_groups)): verbose_proxy_logger.debug( f"User groups {user_groups} matched role '{role.value}' via groups: {role_groups}" ) @@ -799,9 +752,7 @@ def process_sso_jwt_access_token( import jwt try: - access_token_payload = jwt.decode( - access_token_str, options={"verify_signature": False} - ) + access_token_payload = jwt.decode(access_token_str, options={"verify_signature": False}) except jwt.exceptions.DecodeError: verbose_proxy_logger.debug( "Access token is not a valid JWT (possibly an opaque token), skipping JWT-based extraction" @@ -813,41 +764,29 @@ def process_sso_jwt_access_token( if isinstance(result, dict): result_team_ids: Optional[List[str]] = result.get("team_ids", []) if not result_team_ids: - team_ids = sso_jwt_handler.get_team_ids_from_jwt( - access_token_payload - ) + team_ids = sso_jwt_handler.get_team_ids_from_jwt(access_token_payload) result["team_ids"] = team_ids else: result_team_ids = getattr(result, "team_ids", []) if result else [] if not result_team_ids: - team_ids = sso_jwt_handler.get_team_ids_from_jwt( - access_token_payload - ) + team_ids = sso_jwt_handler.get_team_ids_from_jwt(access_token_payload) setattr(result, "team_ids", team_ids) # Extract user role from access token if not already set from UserInfo - existing_role = ( - result.get("user_role") - if isinstance(result, dict) - else getattr(result, "user_role", None) - ) + existing_role = result.get("user_role") if isinstance(result, dict) else getattr(result, "user_role", None) if existing_role is None: user_role: Optional[LitellmUserRoles] = None # Try role_mappings first (group-based role determination) if role_mappings is not None and role_mappings.roles: group_claim = role_mappings.group_claim - user_groups_raw: Any = get_nested_value( - access_token_payload, group_claim - ) + user_groups_raw: Any = get_nested_value(access_token_payload, group_claim) user_groups: List[str] = [] if isinstance(user_groups_raw, list): user_groups = [str(g) for g in user_groups_raw] elif isinstance(user_groups_raw, str): - user_groups = [ - g.strip() for g in user_groups_raw.split(",") if g.strip() - ] + user_groups = [g.strip() for g in user_groups_raw.split(",") if g.strip()] elif user_groups_raw is not None: user_groups = [str(user_groups_raw)] @@ -861,12 +800,8 @@ def process_sso_jwt_access_token( # Fallback: try GENERIC_USER_ROLE_ATTRIBUTE on the access token payload if user_role is None: - generic_user_role_attribute_name = os.getenv( - "GENERIC_USER_ROLE_ATTRIBUTE", "role" - ) - user_role_from_token = get_nested_value( - access_token_payload, generic_user_role_attribute_name - ) + generic_user_role_attribute_name = os.getenv("GENERIC_USER_ROLE_ATTRIBUTE", "role") + user_role_from_token = get_nested_value(access_token_payload, generic_user_role_attribute_name) if user_role_from_token is not None: user_role = get_litellm_user_role(user_role_from_token) verbose_proxy_logger.debug( @@ -878,9 +813,7 @@ def process_sso_jwt_access_token( result["user_role"] = user_role else: setattr(result, "user_role", user_role) - verbose_proxy_logger.debug( - f"Set user_role='{user_role}' from JWT access token" - ) + verbose_proxy_logger.debug(f"Set user_role='{user_role}' from JWT access token") return access_token_payload @@ -921,11 +854,7 @@ async def google_login( return admin_ui_disabled() ####### Check if user is a Enterprise / Premium User ####### - if ( - microsoft_client_id is not None - or google_client_id is not None - or generic_client_id is not None - ): + if microsoft_client_id is not None or google_client_id is not None or generic_client_id is not None: if premium_user is not True: # Check if under 'free SSO user' limit if prisma_client is not None: @@ -1032,26 +961,14 @@ def generic_response_convertor( role_mappings: Optional["RoleMappings"] = None, team_mappings: Optional["TeamMappings"] = None, ) -> CustomOpenID: - generic_user_id_attribute_name = os.getenv( - "GENERIC_USER_ID_ATTRIBUTE", "preferred_username" - ) - generic_user_display_name_attribute_name = os.getenv( - "GENERIC_USER_DISPLAY_NAME_ATTRIBUTE", "sub" - ) - generic_user_email_attribute_name = os.getenv( - "GENERIC_USER_EMAIL_ATTRIBUTE", "email" - ) + generic_user_id_attribute_name = os.getenv("GENERIC_USER_ID_ATTRIBUTE", "preferred_username") + generic_user_display_name_attribute_name = os.getenv("GENERIC_USER_DISPLAY_NAME_ATTRIBUTE", "sub") + generic_user_email_attribute_name = os.getenv("GENERIC_USER_EMAIL_ATTRIBUTE", "email") - generic_user_first_name_attribute_name = os.getenv( - "GENERIC_USER_FIRST_NAME_ATTRIBUTE", "first_name" - ) - generic_user_last_name_attribute_name = os.getenv( - "GENERIC_USER_LAST_NAME_ATTRIBUTE", "last_name" - ) + generic_user_first_name_attribute_name = os.getenv("GENERIC_USER_FIRST_NAME_ATTRIBUTE", "first_name") + generic_user_last_name_attribute_name = os.getenv("GENERIC_USER_LAST_NAME_ATTRIBUTE", "last_name") - generic_provider_attribute_name = os.getenv( - "GENERIC_USER_PROVIDER_ATTRIBUTE", "provider" - ) + generic_provider_attribute_name = os.getenv("GENERIC_USER_PROVIDER_ATTRIBUTE", "provider") generic_user_role_attribute_name = os.getenv("GENERIC_USER_ROLE_ATTRIBUTE", "role") @@ -1118,9 +1035,7 @@ def generic_response_convertor( # Fallback to existing logic if role_mappings not used if user_role is None: - user_role_from_sso = get_nested_value( - response, generic_user_role_attribute_name - ) + user_role_from_sso = get_nested_value(response, generic_user_role_attribute_name) if user_role_from_sso is not None: role = get_litellm_user_role(user_role_from_sso) if role is not None: @@ -1139,12 +1054,8 @@ def generic_response_convertor( return CustomOpenID( id=get_nested_value(response, generic_user_id_attribute_name), - display_name=get_nested_value( - response, generic_user_display_name_attribute_name - ), - email=normalize_email( - get_nested_value(response, generic_user_email_attribute_name) - ), + display_name=get_nested_value(response, generic_user_display_name_attribute_name), + email=normalize_email(get_nested_value(response, generic_user_email_attribute_name)), first_name=get_nested_value(response, generic_user_first_name_attribute_name), last_name=get_nested_value(response, generic_user_last_name_attribute_name), provider=get_nested_value(response, generic_provider_attribute_name), @@ -1163,9 +1074,7 @@ def _setup_generic_sso_env_vars( generic_authorization_endpoint = os.getenv("GENERIC_AUTHORIZATION_ENDPOINT", None) generic_token_endpoint = os.getenv("GENERIC_TOKEN_ENDPOINT", None) generic_userinfo_endpoint = os.getenv("GENERIC_USERINFO_ENDPOINT", None) - generic_include_client_id = ( - os.getenv("GENERIC_INCLUDE_CLIENT_ID", "false").lower() == "true" - ) + generic_include_client_id = os.getenv("GENERIC_INCLUDE_CLIENT_ID", "false").lower() == "true" # Validate required environment variables if generic_client_secret is None: @@ -1200,9 +1109,7 @@ def _setup_generic_sso_env_vars( verbose_proxy_logger.debug( f"authorization_endpoint: {generic_authorization_endpoint}\ntoken_endpoint: {generic_token_endpoint}\nuserinfo_endpoint: {generic_userinfo_endpoint}" ) - verbose_proxy_logger.debug( - f"GENERIC_REDIRECT_URI: {redirect_url}\nGENERIC_CLIENT_ID: {generic_client_id}\n" - ) + verbose_proxy_logger.debug(f"GENERIC_REDIRECT_URI: {redirect_url}\nGENERIC_CLIENT_ID: {generic_client_id}\n") return ( generic_client_secret, @@ -1220,13 +1127,9 @@ async def _setup_team_mappings() -> Optional["TeamMappings"]: try: from litellm.proxy.utils import get_prisma_client_or_throw - prisma_client = get_prisma_client_or_throw( - "Prisma client is None, connect a database to your proxy" - ) + prisma_client = get_prisma_client_or_throw("Prisma client is None, connect a database to your proxy") - sso_db_record = await SSOConfigRepository(prisma_client).table.find_unique( - where={"id": "sso_config"} - ) + sso_db_record = await SSOConfigRepository(prisma_client).table.find_unique(where={"id": "sso_config"}) if sso_db_record and sso_db_record.sso_settings: sso_settings_dict = dict(sso_db_record.sso_settings) @@ -1258,13 +1161,9 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]: try: from litellm.proxy.utils import get_prisma_client_or_throw - prisma_client = get_prisma_client_or_throw( - "Prisma client is None, connect a database to your proxy" - ) + prisma_client = get_prisma_client_or_throw("Prisma client is None, connect a database to your proxy") - sso_db_record = await SSOConfigRepository(prisma_client).table.find_unique( - where={"id": "sso_config"} - ) + sso_db_record = await SSOConfigRepository(prisma_client).table.find_unique(where={"id": "sso_config"}) if sso_db_record and sso_db_record.sso_settings: sso_settings_dict = dict(sso_db_record.sso_settings) @@ -1279,31 +1178,21 @@ async def _setup_role_mappings() -> Optional["RoleMappings"]: role_mappings = role_mappings_data if role_mappings: - verbose_proxy_logger.debug( - f"Loaded role_mappings for provider '{role_mappings.provider}'" - ) + verbose_proxy_logger.debug(f"Loaded role_mappings for provider '{role_mappings.provider}'") except Exception as e: verbose_proxy_logger.debug( f"Could not load role_mappings from database: {e}. Continuing with existing role logic." ) generic_role_mappings = os.getenv("GENERIC_ROLE_MAPPINGS_ROLES", None) - generic_role_mappings_group_claim = os.getenv( - "GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", None - ) - generic_role_mappings_default_role = os.getenv( - "GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE", None - ) + generic_role_mappings_group_claim = os.getenv("GENERIC_ROLE_MAPPINGS_GROUP_CLAIM", None) + generic_role_mappings_default_role = os.getenv("GENERIC_ROLE_MAPPINGS_DEFAULT_ROLE", None) if generic_role_mappings is not None: - verbose_proxy_logger.debug( - "Found role_mappings for generic provider in environment variables" - ) + verbose_proxy_logger.debug("Found role_mappings for generic provider in environment variables") import ast try: - generic_user_role_mappings_data: Dict[LitellmUserRoles, List[str]] = ( - ast.literal_eval(generic_role_mappings) - ) + generic_user_role_mappings_data: Dict[LitellmUserRoles, List[str]] = ast.literal_eval(generic_role_mappings) if isinstance(generic_user_role_mappings_data, dict): from litellm.types.proxy.management_endpoints.ui_sso import RoleMappings @@ -1353,13 +1242,10 @@ def _handle_generic_sso_error( # 1. The error mentions PKCE/code verifier, AND # 2. PKCE is not currently configured (GENERIC_CLIENT_USE_PKCE != true) pkce_configured = os.getenv("GENERIC_CLIENT_USE_PKCE", "false").lower() == "true" - if not pkce_configured and ( - "PKCE" in error_message or "code verifier" in error_message.lower() - ): - is_okta = ( - generic_authorization_endpoint - and "okta" in generic_authorization_endpoint.lower() - ) or (generic_token_endpoint and "okta" in generic_token_endpoint.lower()) + if not pkce_configured and ("PKCE" in error_message or "code verifier" in error_message.lower()): + is_okta = (generic_authorization_endpoint and "okta" in generic_authorization_endpoint.lower()) or ( + generic_token_endpoint and "okta" in generic_token_endpoint.lower() + ) provider_name = "Okta" if is_okta else "Your OAuth provider" detailed_message = ( @@ -1401,14 +1287,10 @@ def _handle_generic_sso_error( async def get_generic_sso_response( request: Request, jwt_handler: JWTHandler, - sso_jwt_handler: Optional[ - JWTHandler - ], # sso specific jwt handler - used for restricted sso group access control + sso_jwt_handler: Optional[JWTHandler], # sso specific jwt handler - used for restricted sso group access control generic_client_id: str, redirect_url: str, -) -> Tuple[ - Union[OpenID, dict], Optional[dict], Optional[dict] -]: # (result, received_response, access_token_payload) +) -> Tuple[Union[OpenID, dict], Optional[dict], Optional[dict]]: # (result, received_response, access_token_payload) # make generic sso provider from fastapi_sso.sso.base import DiscoveryDocument from fastapi_sso.sso.generic import create_provider @@ -1460,17 +1342,13 @@ async def get_generic_sso_response( verbose_proxy_logger.debug("calling generic_sso.verify_and_process") additional_generic_sso_headers_dict = _parse_generic_sso_headers() - code_verifier: Optional[str] = ( - None # assigned inside try; initialized for type tracking - ) + code_verifier: Optional[str] = None # assigned inside try; initialized for type tracking access_token_payload: Optional[dict] = None # decoded JWT access token claims try: - token_exchange_params = ( - await SSOAuthenticationHandler.prepare_token_exchange_parameters( - request=request, - generic_include_client_id=generic_include_client_id, - ) + token_exchange_params = await SSOAuthenticationHandler.prepare_token_exchange_parameters( + request=request, + generic_include_client_id=generic_include_client_id, ) # Extract code_verifier (and the cache key for deferred deletion) before calling fastapi-sso @@ -1493,16 +1371,9 @@ async def get_generic_sso_response( # code (Login-CSRF / token theft). url_state = request.query_params.get("state") cookie_state = request.cookies.get("litellm_oauth_state") - if ( - not url_state - or not cookie_state - or not secrets.compare_digest(url_state, cookie_state) - ): + if not url_state or not cookie_state or not secrets.compare_digest(url_state, cookie_state): raise ProxyException( - message=( - "Invalid OAuth state parameter — does not match " - "the browser-bound state cookie." - ), + message=("Invalid OAuth state parameter — does not match the browser-bound state cookie."), type=ProxyErrorTypes.auth_error, param="state", code=status.HTTP_400_BAD_REQUEST, @@ -1557,11 +1428,7 @@ async def get_generic_sso_response( # must not be exposed to callers. # Assign directly rather than relying on nonlocal mutation so that Pyright # can track that received_response is non-None from this point on. - received_response = { - k: v - for k, v in combined_response.items() - if k not in _OAUTH_TOKEN_FIELDS - } + received_response = {k: v for k, v in combined_response.items() if k not in _OAUTH_TOKEN_FIELDS} # In the PKCE path verify_and_process is skipped, so generic_sso.access_token # is never set. Read the token directly from the exchange response instead so # process_sso_jwt_access_token can extract JWT-embedded roles/teams. @@ -1608,14 +1475,10 @@ async def create_team_member_add_task(team_id, user_info): user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN), ) except Exception as e: - verbose_proxy_logger.debug( - f"[Non-Blocking] Error trying to add sso user to db: {e}" - ) + verbose_proxy_logger.debug(f"[Non-Blocking] Error trying to add sso user to db: {e}") -async def add_missing_team_member( - user_info: Union[NewUserResponse, LiteLLM_UserTable], sso_teams: List[str] -): +async def add_missing_team_member(user_info: Union[NewUserResponse, LiteLLM_UserTable], sso_teams: List[str]): """ - Get missing teams (diff b/w user_info.team_ids and sso_teams) - Add missing user to missing teams @@ -1625,26 +1488,19 @@ async def add_missing_team_member( missing_teams = set(sso_teams) - set(user_teams) missing_teams_list = list(missing_teams) tasks = [] - tasks = [ - create_team_member_add_task(team_id, user_info) - for team_id in missing_teams_list - ] + tasks = [create_team_member_add_task(team_id, user_info) for team_id in missing_teams_list] try: await asyncio.gather(*tasks) except Exception as e: - verbose_proxy_logger.debug( - f"[Non-Blocking] Error trying to add sso user to db: {e}" - ) + verbose_proxy_logger.debug(f"[Non-Blocking] Error trying to add sso user to db: {e}") def get_disabled_non_admin_personal_key_creation(): key_generation_settings = litellm.key_generation_settings if key_generation_settings is None: return False - personal_key_generation = ( - key_generation_settings.get("personal_key_generation") or {} - ) + personal_key_generation = key_generation_settings.get("personal_key_generation") or {} allowed_user_roles = personal_key_generation.get("allowed_user_roles") or [] return bool("proxy_admin" in allowed_user_roles) @@ -1697,9 +1553,7 @@ async def get_user_info_from_db( potential_user_ids.append(_id) user_email = normalize_email( - getattr(result, "email", None) - if not isinstance(result, dict) - else result.get("email", None) + getattr(result, "email", None) if not isinstance(result, dict) else result.get("email", None) ) user_info: Optional[Union[LiteLLM_UserTable, NewUserResponse]] = None @@ -1737,9 +1591,7 @@ async def get_user_info_from_db( except ProxyException: raise except Exception as e: - verbose_proxy_logger.exception( - f"[Non-Blocking] Error trying to add sso user to db: {e}" - ) + verbose_proxy_logger.exception(f"[Non-Blocking] Error trying to add sso user to db: {e}") return None @@ -1780,16 +1632,12 @@ def _build_sso_user_update_data( sso_role = getattr(result, "user_role", None) if sso_role is not None: # Convert enum to string if needed - sso_role_str = ( - sso_role.value if isinstance(sso_role, LitellmUserRoles) else sso_role - ) + sso_role_str = sso_role.value if isinstance(sso_role, LitellmUserRoles) else sso_role # Only include if it's a valid LiteLLM role if _should_use_role_from_sso_response(sso_role_str): update_data["user_role"] = sso_role_str - verbose_proxy_logger.info( - f"Updating user {user_id} role from SSO: {sso_role_str}" - ) + verbose_proxy_logger.info(f"Updating user {user_id} role from SSO: {sso_role_str}") return update_data @@ -1821,9 +1669,7 @@ async def _sync_user_role_from_jwt_role_map( if mapped_role is None: return - verbose_proxy_logger.info( - f"SSO jwt_litellm_role_map matched role: {mapped_role.value}" - ) + verbose_proxy_logger.info(f"SSO jwt_litellm_role_map matched role: {mapped_role.value}") # Update user_defined_values so downstream code uses the mapped role if user_defined_values is not None: @@ -1859,18 +1705,12 @@ def apply_user_info_values_to_sso_user_defined_values( if _should_use_role_from_sso_response(sso_role): # SSO provided a valid role, keep it and log that we're using it - verbose_proxy_logger.info( - f"Using SSO role: {sso_role} (DB role was: {db_role})" - ) + verbose_proxy_logger.info(f"Using SSO role: {sso_role} (DB role was: {db_role})") else: # SSO didn't provide a valid role, fall back to DB role or default if user_info is None or user_info.user_role is None: - user_defined_values["user_role"] = ( - LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value - ) - verbose_proxy_logger.debug( - "No SSO or DB role found, using default: INTERNAL_USER_VIEW_ONLY" - ) + user_defined_values["user_role"] = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value + verbose_proxy_logger.debug("No SSO or DB role found, using default: INTERNAL_USER_VIEW_ONLY") else: user_defined_values["user_role"] = user_info.user_role verbose_proxy_logger.debug(f"Using DB role: {user_info.user_role}") @@ -1882,9 +1722,7 @@ def apply_user_info_values_to_sso_user_defined_values( return user_defined_values -async def check_and_update_if_proxy_admin_id( - user_role: str, user_id: str, prisma_client: Optional[PrismaClient] -): +async def check_and_update_if_proxy_admin_id(user_role: str, user_id: str, prisma_client: Optional[PrismaClient]): """ - Check if user role in DB is admin - If not, update user role in DB to admin role @@ -1923,9 +1761,7 @@ async def auth_callback(request: Request, state: Optional[str] = None): ) if prisma_client is None: - raise HTTPException( - status_code=500, detail=CommonProxyErrors.db_not_connected_error.value - ) + raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) sso_jwt_handler: Optional[JWTHandler] = None ui_access_mode = general_settings.get("ui_access_mode", None) @@ -1935,9 +1771,7 @@ async def auth_callback(request: Request, state: Optional[str] = None): prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, litellm_jwtauth=LiteLLM_JWTAuth( - team_ids_jwt_field=general_settings.get("ui_access_mode", {}).get( - "sso_group_jwt_field", None - ), + team_ids_jwt_field=general_settings.get("ui_access_mode", {}).get("sso_group_jwt_field", None), ), leeway=0, ) @@ -1955,9 +1789,7 @@ async def auth_callback(request: Request, state: Optional[str] = None): param="master_key", code=status.HTTP_500_INTERNAL_SERVER_ERROR, ) - redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso( - request=request, sso_callback_route="sso/callback" - ) + redirect_url = SSOAuthenticationHandler.get_redirect_url_for_sso(request=request, sso_callback_route="sso/callback") verbose_proxy_logger.info(f"Redirecting to {redirect_url}") result = None @@ -2054,9 +1886,7 @@ async def _fetch_cli_sso_team_details( team_details: List[Dict[str, Any]] = [] try: if teams: - prisma_teams = await TeamRepository(prisma_client).table.find_many( - where={"team_id": {"in": teams}} - ) + prisma_teams = await TeamRepository(prisma_client).table.find_many(where={"team_id": {"in": teams}}) for team_row in prisma_teams: team_dict = team_row.model_dump() team_details.append( @@ -2066,9 +1896,7 @@ async def _fetch_cli_sso_team_details( } ) except Exception as e: - verbose_proxy_logger.error( - f"Error fetching team details for CLI SSO session: {e}" - ) + verbose_proxy_logger.error(f"Error fetching team details for CLI SSO session: {e}") return team_details @@ -2099,21 +1927,15 @@ async def _complete_cli_sso_callback_session( alternate_user_id=user_id, ) if user_info is None: - raise HTTPException( - status_code=500, detail="Failed to retrieve user information from SSO" - ) + raise HTTPException(status_code=500, detail="Failed to retrieve user information from SSO") if not user_info.user_id: - raise HTTPException( - status_code=500, detail="Failed to retrieve user information from SSO" - ) + raise HTTPException(status_code=500, detail="Failed to retrieve user information from SSO") teams: List[str] = [] if hasattr(user_info, "teams") and user_info.teams: teams = user_info.teams if isinstance(user_info.teams, list) else [] - team_details = await _fetch_cli_sso_team_details( - prisma_client=prisma_client, teams=teams - ) + team_details = await _fetch_cli_sso_team_details(prisma_client=prisma_client, teams=teams) attribution_metadata = build_cli_sso_attribution_metadata(result=result) if attribution_metadata: await _persist_cli_sso_user_metadata( @@ -2173,9 +1995,7 @@ async def cli_sso_callback( flow = _get_cli_sso_flow_or_raise(login_id=key, cache=user_api_key_cache) if prisma_client is None: - raise HTTPException( - status_code=500, detail=CommonProxyErrors.db_not_connected_error.value - ) + raise HTTPException(status_code=500, detail=CommonProxyErrors.db_not_connected_error.value) if result is None: raise HTTPException( @@ -2187,11 +2007,9 @@ async def cli_sso_callback( result_non_none: Union[OpenID, dict] = cast(Union[OpenID, dict], result) try: - parsed_openid_result = ( - SSOAuthenticationHandler._get_user_email_and_id_from_result( - result=result_non_none, - generic_client_id=os.getenv("GENERIC_CLIENT_ID", None), - ) + parsed_openid_result = SSOAuthenticationHandler._get_user_email_and_id_from_result( + result=result_non_none, + generic_client_id=os.getenv("GENERIC_CLIENT_ID", None), ) verbose_proxy_logger.debug(f"parsed_openid_result: {parsed_openid_result}") user_defined_values = await _build_cli_sso_user_defined_values( @@ -2223,9 +2041,7 @@ async def cli_sso_callback( raise except Exception as e: verbose_proxy_logger.error(f"Error with CLI SSO callback: {e}") - raise HTTPException( - status_code=500, detail=f"Failed to process CLI SSO: {str(e)}" - ) + raise HTTPException(status_code=500, detail=f"Failed to process CLI SSO: {str(e)}") @router.get("/sso/cli/poll/{key_id}", tags=["experimental"], include_in_schema=False) @@ -2250,9 +2066,7 @@ async def cli_poll_key( try: flow = _get_cli_sso_flow_or_raise(login_id=key_id, cache=user_api_key_cache) - if not _verify_cli_sso_poll_secret( - flow=flow, poll_secret=x_litellm_cli_poll_secret - ): + if not _verify_cli_sso_poll_secret(flow=flow, poll_secret=x_litellm_cli_poll_secret): raise HTTPException(status_code=403, detail="Invalid CLI polling secret") if not flow.get("sso_complete") or not flow.get("user_code_verified"): @@ -2274,18 +2088,14 @@ async def cli_poll_key( # clients we return rich team details (id + alias); older clients # can continue to rely on the simple "teams" list. if team_id is None and len(user_teams) > 1: - verbose_proxy_logger.info( - f"Returning teams list for user {user_id} to select from: {user_teams}" - ) + verbose_proxy_logger.info(f"Returning teams list for user {user_id} to select from: {user_teams}") # Best-effort construction of team_details if it wasn't # already cached for some reason. team_details_response: Optional[List[Dict[str, Any]]] = None if isinstance(user_team_details, list) and user_team_details: team_details_response = user_team_details elif user_teams: - team_details_response = [ - {"team_id": t, "team_alias": None} for t in user_teams - ] + team_details_response = [{"team_id": t, "team_alias": None} for t in user_teams] poll_response: Dict[str, Any] = { "status": "ready", "user_id": user_id, @@ -2293,9 +2103,7 @@ async def cli_poll_key( "team_details": team_details_response, "requires_team_selection": True, } - attribution_metadata = _cli_poll_attribution_metadata_from_session( - session_data - ) + attribution_metadata = _cli_poll_attribution_metadata_from_session(session_data) if attribution_metadata: poll_response["attribution_metadata"] = attribution_metadata return poll_response @@ -2314,11 +2122,7 @@ async def cli_poll_key( team_alias = None if team_id and isinstance(user_team_details, list): team_alias = next( - ( - team.get("team_alias") - for team in user_team_details - if team.get("team_id") == team_id - ), + (team.get("team_alias") for team in user_team_details if team.get("team_id") == team_id), None, ) @@ -2327,21 +2131,49 @@ async def cli_poll_key( user_id=user_id, user_role=session_data["user_role"], models=session_data.get("models", []), - max_budget=litellm.max_ui_session_budget, ) + # Resolve effective session budget cap: + # - If user has their own budget, let it apply (no cap override) + # - If neither user nor team has a budget, cap at max_ui_session_budget + from litellm.proxy.auth.auth_checks import get_team_object, get_user_object + from litellm.proxy.proxy_server import prisma_client + + resolved_max_budget: Optional[float] = None + try: + db_user = await get_user_object( + user_id=user_id, + prisma_client=prisma_client, + user_id_upsert=False, + user_api_key_cache=user_api_key_cache, + ) + user_has_budget = db_user is not None and db_user.max_budget is not None + except Exception: + user_has_budget = True # assume budget exists if lookup fails + + if not user_has_budget and team_id: + try: + db_team = await get_team_object( + team_id=team_id, + prisma_client=prisma_client, + user_api_key_cache=user_api_key_cache, + ) + team_has_budget = db_team is not None and db_team.max_budget is not None + if not team_has_budget: + resolved_max_budget = litellm.max_ui_session_budget + except Exception: + pass # team lookup failed: don't apply fallback cap + # Generate CLI JWT on-demand (expiration configurable via LITELLM_CLI_JWT_EXPIRATION_HOURS) # Pass selected team_id to ensure JWT has correct team jwt_token = ExperimentalUIJWTToken.get_cli_jwt_auth_token( - user_info=user_info, team_id=team_id, team_alias=team_alias + user_info=user_info, team_id=team_id, team_alias=team_alias, max_budget=resolved_max_budget ) # Delete cache entry (single-use) user_api_key_cache.delete_cache(key=_get_cli_sso_flow_cache_key(key_id)) - verbose_proxy_logger.info( - f"CLI JWT generated for user: {user_id}, team: {team_id}" - ) + verbose_proxy_logger.info(f"CLI JWT generated for user: {user_id}, team: {team_id}") poll_response = { "status": "ready", "key": jwt_token, @@ -2352,9 +2184,7 @@ async def cli_poll_key( # present nicer information if needed. "team_details": user_team_details, } - attribution_metadata = _cli_poll_attribution_metadata_from_session( - session_data - ) + attribution_metadata = _cli_poll_attribution_metadata_from_session(session_data) if attribution_metadata: poll_response["attribution_metadata"] = attribution_metadata return poll_response @@ -2365,9 +2195,7 @@ async def cli_poll_key( raise except Exception as e: verbose_proxy_logger.error(f"Error polling for CLI JWT: {e}") - raise HTTPException( - status_code=500, detail=f"Error checking session status: {str(e)}" - ) + raise HTTPException(status_code=500, detail=f"Error checking session status: {str(e)}") async def _enforce_free_sso_user_limit( @@ -2413,9 +2241,7 @@ async def insert_sso_user( Returns: Tuple[str, str]: User ID and User Role """ - verbose_proxy_logger.debug( - f"Inserting SSO user into DB. User values: {user_defined_values}" - ) + verbose_proxy_logger.debug(f"Inserting SSO user into DB. User values: {user_defined_values}") if prisma_client is not None: from litellm.proxy.proxy_server import premium_user as _premium_user @@ -2443,9 +2269,7 @@ async def insert_sso_user( preserved_role = sso_role user_defined_values.update(litellm.default_internal_user_params) # type: ignore user_defined_values["user_role"] = preserved_role # Restore preserved role - verbose_proxy_logger.debug( - f"Preserved SSO-extracted role '{preserved_role}'" - ) + verbose_proxy_logger.debug(f"Preserved SSO-extracted role '{preserved_role}'") else: # SSO didn't provide a valid role, apply all defaults including role user_defined_values.update(litellm.default_internal_user_params) # type: ignore @@ -2455,9 +2279,7 @@ async def insert_sso_user( if user_defined_values.get("max_budget") is None: user_defined_values["max_budget"] = litellm.max_internal_user_budget if user_defined_values.get("budget_duration") is None: - user_defined_values["budget_duration"] = ( - litellm.internal_user_budget_duration - ) + user_defined_values["budget_duration"] = litellm.internal_user_budget_duration if user_defined_values["user_role"] is None: user_defined_values["user_role"] = LitellmUserRoles.INTERNAL_USER_VIEW_ONLY @@ -2473,9 +2295,7 @@ async def insert_sso_user( ) if result_openid and hasattr(result_openid, "provider"): - new_user_request.metadata = { - "auth_provider": getattr(result_openid, "provider") - } + new_user_request.metadata = {"auth_provider": getattr(result_openid, "provider")} response = await new_user( data=new_user_request, @@ -2499,8 +2319,7 @@ async def get_ui_settings(request: Request): _api_doc_base_url = os.getenv("LITELLM_UI_API_DOC_BASE_URL", None) _is_sso_enabled = _has_user_setup_sso() disable_expensive_db_queries = ( - proxy_state.get_proxy_state_variable("spend_logs_row_count") - > MAX_SPENDLOG_ROWS_TO_QUERY + proxy_state.get_proxy_state_variable("spend_logs_row_count") > MAX_SPENDLOG_ROWS_TO_QUERY ) default_team_disabled = general_settings.get("default_team_disabled", False) if "PROXY_DEFAULT_TEAM_DISABLED" in os.environ: @@ -2513,9 +2332,7 @@ async def get_ui_settings(request: Request): "LITELLM_UI_API_DOC_BASE_URL": _api_doc_base_url, "DEFAULT_TEAM_DISABLED": default_team_disabled, "SSO_ENABLED": _is_sso_enabled, - "NUM_SPEND_LOGS_ROWS": proxy_state.get_proxy_state_variable( - "spend_logs_row_count" - ), + "NUM_SPEND_LOGS_ROWS": proxy_state.get_proxy_state_variable("spend_logs_row_count"), "DISABLE_EXPENSIVE_DB_QUERIES": disable_expensive_db_queries, } @@ -2569,9 +2386,7 @@ async def sso_readiness(): elif configured_provider == "generic": generic_client_secret = os.getenv("GENERIC_CLIENT_SECRET", None) - generic_authorization_endpoint = os.getenv( - "GENERIC_AUTHORIZATION_ENDPOINT", None - ) + generic_authorization_endpoint = os.getenv("GENERIC_AUTHORIZATION_ENDPOINT", None) generic_token_endpoint = os.getenv("GENERIC_TOKEN_ENDPOINT", None) generic_userinfo_endpoint = os.getenv("GENERIC_USERINFO_ENDPOINT", None) if generic_client_secret is None: @@ -2711,12 +2526,8 @@ class SSOAuthenticationHandler: from fastapi_sso.sso.generic import create_provider generic_client_secret = os.getenv("GENERIC_CLIENT_SECRET", None) - generic_scope = os.getenv("GENERIC_SCOPE", "openid email profile").split( - " " - ) - generic_authorization_endpoint = os.getenv( - "GENERIC_AUTHORIZATION_ENDPOINT", None - ) + generic_scope = os.getenv("GENERIC_SCOPE", "openid email profile").split(" ") + generic_authorization_endpoint = os.getenv("GENERIC_AUTHORIZATION_ENDPOINT", None) generic_token_endpoint = os.getenv("GENERIC_TOKEN_ENDPOINT", None) generic_userinfo_endpoint = os.getenv("GENERIC_USERINFO_ENDPOINT", None) if generic_client_secret is None: @@ -2836,9 +2647,7 @@ class SSOAuthenticationHandler: value={"code_verifier": code_verifier}, ttl=600, ) - verbose_proxy_logger.debug( - "PKCE code_verifier stored in cache (TTL: 600s)" - ) + verbose_proxy_logger.debug("PKCE code_verifier stored in cache (TTL: 600s)") # Add PKCE parameters to the authorization URL if pkce_params: @@ -2950,11 +2759,7 @@ class SSOAuthenticationHandler: microsoft_client_id: Optional[str] = None, generic_client_id: Optional[str] = None, ) -> bool: - if ( - google_client_id is not None - or microsoft_client_id is not None - or generic_client_id is not None - ): + if google_client_id is not None or microsoft_client_id is not None or generic_client_id is not None: return True return False @@ -3003,13 +2808,9 @@ class SSOAuthenticationHandler: user_id=user_id, ) - await UserRepository(prisma_client).table.update_many( - where={"user_id": user_id}, data=update_data - ) + await UserRepository(prisma_client).table.update_many(where={"user_id": user_id}, data=update_data) else: - verbose_proxy_logger.info( - "user not in DB, inserting user into LiteLLM DB" - ) + verbose_proxy_logger.info("user not in DB, inserting user into LiteLLM DB") # user not in DB, insert User into LiteLLM DB user_info = await insert_sso_user( result_openid=result, @@ -3020,9 +2821,7 @@ class SSOAuthenticationHandler: except ProxyException: raise except Exception as e: - verbose_proxy_logger.exception( - f"Error upserting SSO user into LiteLLM DB: {e}" - ) + verbose_proxy_logger.exception(f"Error upserting SSO user into LiteLLM DB: {e}") return user_info @staticmethod @@ -3037,9 +2836,7 @@ class SSOAuthenticationHandler: The `team_ids` field is populated by litellm after processing the SSO response """ if user_info is None: - verbose_proxy_logger.debug( - "User not found in LiteLLM DB, skipping team member addition" - ) + verbose_proxy_logger.debug("User not found in LiteLLM DB, skipping team member addition") return sso_teams = getattr(result, "team_ids", []) await add_missing_team_member(user_info=user_info, sso_teams=sso_teams) @@ -3061,9 +2858,7 @@ class SSOAuthenticationHandler: - if result.team_ids is a list, return True if the restricted_sso_group is in the list, otherwise return False """ - ui_access_mode = cast( - Optional[Union[Dict, str]], general_settings.get("ui_access_mode") - ) + ui_access_mode = cast(Optional[Union[Dict, str]], general_settings.get("ui_access_mode")) if ui_access_mode is None: return True @@ -3108,16 +2903,12 @@ class SSOAuthenticationHandler: code=status.HTTP_500_INTERNAL_SERVER_ERROR, ) try: - team_obj = await TeamRepository(prisma_client).table.find_first( - where={"team_id": litellm_team_id} - ) + team_obj = await TeamRepository(prisma_client).table.find_first(where={"team_id": litellm_team_id}) verbose_proxy_logger.debug(f"Team object: {team_obj}") # only create a new team if it doesn't exist if team_obj: - verbose_proxy_logger.debug( - f"Team already exists: {litellm_team_id} - {litellm_team_name}" - ) + verbose_proxy_logger.debug(f"Team already exists: {litellm_team_id} - {litellm_team_name}") return team_request: NewTeamRequest = NewTeamRequest( @@ -3206,9 +2997,7 @@ class SSOAuthenticationHandler: Gets the user email and id from the OpenID result after validating the email domain """ user_email: Optional[str] = normalize_email(getattr(result, "email", None)) - user_id: Optional[str] = ( - getattr(result, "id", None) if result is not None else None - ) + user_id: Optional[str] = getattr(result, "id", None) if result is not None else None user_role: Optional[str] = None if user_email is not None and os.getenv("ALLOWED_EMAIL_DOMAINS") is not None: @@ -3229,20 +3018,12 @@ class SSOAuthenticationHandler: _user_role = getattr(result, "user_role", None) if _user_role is not None: # Convert enum to string if needed - user_role = ( - _user_role.value - if isinstance(_user_role, LitellmUserRoles) - else _user_role - ) - verbose_proxy_logger.debug( - f"Extracted user_role from SSO result: {user_role}" - ) + user_role = _user_role.value if isinstance(_user_role, LitellmUserRoles) else _user_role + verbose_proxy_logger.debug(f"Extracted user_role from SSO result: {user_role}") # generic client id - override with custom attribute name if specified if generic_client_id is not None and result is not None: - generic_user_role_attribute_name = os.getenv( - "GENERIC_USER_ROLE_ATTRIBUTE", "role" - ) + generic_user_role_attribute_name = os.getenv("GENERIC_USER_ROLE_ATTRIBUTE", "role") user_id = getattr(result, "id", None) user_email = normalize_email(getattr(result, "email", None)) if user_role is None: @@ -3250,9 +3031,7 @@ class SSOAuthenticationHandler: if _role_from_attr is not None: # Convert enum to string if needed user_role = ( - _role_from_attr.value - if isinstance(_role_from_attr, LitellmUserRoles) - else _role_from_attr + _role_from_attr.value if isinstance(_role_from_attr, LitellmUserRoles) else _role_from_attr ) if user_id is None and result is not None: @@ -3295,15 +3074,11 @@ class SSOAuthenticationHandler: from litellm.proxy.utils import get_prisma_client_or_throw from litellm.types.proxy.ui_sso import ReturnedUITokenObject - prisma_client = get_prisma_client_or_throw( - "Prisma client is None, connect a database to your proxy" - ) + prisma_client = get_prisma_client_or_throw("Prisma client is None, connect a database to your proxy") # User is Authe'd in - generate key for the UI to access Proxy - parsed_openid_result = ( - SSOAuthenticationHandler._get_user_email_and_id_from_result( - result=result, generic_client_id=generic_client_id - ) + parsed_openid_result = SSOAuthenticationHandler._get_user_email_and_id_from_result( + result=result, generic_client_id=generic_client_id ) user_email = parsed_openid_result.get("user_email") user_id = parsed_openid_result.get("user_id") @@ -3382,9 +3157,7 @@ class SSOAuthenticationHandler: "Unable to map user identity to known values. 'user_defined_values' is None. File an issue - https://github.com/BerriAI/litellm/issues" ) - verbose_proxy_logger.info( - f"user_defined_values for creating ui key: {user_defined_values}" - ) + verbose_proxy_logger.info(f"user_defined_values for creating ui key: {user_defined_values}") default_ui_key_values.update(user_defined_values) default_ui_key_values["request_type"] = "key" @@ -3396,18 +3169,13 @@ class SSOAuthenticationHandler: key = response["token"] # type: ignore user_id = response["user_id"] # type: ignore - user_role = ( - user_defined_values["user_role"] - or LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value - ) + user_role = user_defined_values["user_role"] or LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value if user_id and isinstance(user_id, str): user_role = await check_and_update_if_proxy_admin_id( user_role=user_role, user_id=user_id, prisma_client=prisma_client ) - verbose_proxy_logger.debug( - f"user_role: {user_role}; ui_access_mode: {ui_access_mode}" - ) + verbose_proxy_logger.debug(f"user_role: {user_role}; ui_access_mode: {ui_access_mode}") ## CHECK IF ROLE ALLOWED TO USE PROXY ## is_admin_only_access = check_is_admin_only_access(ui_access_mode or {}) if is_admin_only_access: @@ -3420,19 +3188,12 @@ class SSOAuthenticationHandler: }, ) - disabled_non_admin_personal_key_creation = ( - get_disabled_non_admin_personal_key_creation() - ) - litellm_dashboard_ui = get_custom_url( - request_base_url=str(request.base_url), route="ui/" - ) + disabled_non_admin_personal_key_creation = get_disabled_non_admin_personal_key_creation() + litellm_dashboard_ui = get_custom_url(request_base_url=str(request.base_url), route="ui/") if get_secret_bool("EXPERIMENTAL_UI_LOGIN"): _user_info: Optional[LiteLLM_UserTable] = None - if ( - user_defined_values is not None - and user_defined_values["user_id"] is not None - ): + if user_defined_values is not None and user_defined_values["user_id"] is not None: _user_info = LiteLLM_UserTable( user_id=user_defined_values["user_id"], user_role=user_defined_values["user_role"] or user_role, @@ -3442,14 +3203,10 @@ class SSOAuthenticationHandler: if _user_info is None: raise HTTPException( status_code=401, - detail={ - "error": "User Information is required for experimental UI login" - }, + detail={"error": "User Information is required for experimental UI login"}, ) - key = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token( - _user_info - ) + key = ExperimentalUIJWTToken.get_experimental_ui_login_jwt_auth_token(_user_info) returned_ui_token_object = ReturnedUITokenObject( user_id=cast(str, user_id), @@ -3458,9 +3215,7 @@ class SSOAuthenticationHandler: user_role=user_role or LitellmUserRoles.INTERNAL_USER_VIEW_ONLY.value, login_method="sso", premium_user=premium_user, - auth_header_name=general_settings.get( - "litellm_key_header_name", "Authorization" - ), + auth_header_name=general_settings.get("litellm_key_header_name", "Authorization"), disabled_non_admin_personal_key_creation=disabled_non_admin_personal_key_creation, server_root_path=get_server_root_path(), ) @@ -3474,28 +3229,18 @@ class SSOAuthenticationHandler: # Control-plane cross-origin: store JWT behind a single-use opaque # code (60s TTL) so the token never appears in browser history / logs. # The control plane redeems it via POST /v3/login/exchange. - if return_to is not None and SSOAuthenticationHandler._validate_return_to( - return_to - ): + if return_to is not None and SSOAuthenticationHandler._validate_return_to(return_to): code = secrets.token_urlsafe(32) cache_key = f"login_code:{code}" cache_value = {"token": jwt_token, "redirect_url": return_to} if redis_usage_cache is not None: - await redis_usage_cache.async_set_cache( - key=cache_key, value=cache_value, ttl=60 - ) + await redis_usage_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60) else: - await user_api_key_cache.async_set_cache( - key=cache_key, value=cache_value, ttl=60 - ) + await user_api_key_cache.async_set_cache(key=cache_key, value=cache_value, ttl=60) separator = "&" if "?" in return_to else "?" - redirect_url = ( - return_to + separator + urlencode({"login": "success", "code": code}) - ) - verbose_proxy_logger.info( - "Cross-origin SSO: redirecting to control plane with login code" - ) + redirect_url = return_to + separator + urlencode({"login": "success", "code": code}) + verbose_proxy_logger.info("Cross-origin SSO: redirecting to control plane with login code") redirect_response = RedirectResponse(url=redirect_url, status_code=303) redirect_response.delete_cookie("litellm_cp_return_to") return redirect_response @@ -3570,15 +3315,12 @@ class SSOAuthenticationHandler: state, ) else: - verbose_proxy_logger.debug( - "PKCE code_verifier retrieved from cache" - ) + verbose_proxy_logger.debug("PKCE code_verifier retrieved from cache") elif isinstance(cached_data, str): # Handle legacy format (plain string) for backward compatibility code_verifier = cached_data verbose_proxy_logger.warning( - "Retrieved code_verifier in legacy plain-string format. " - "Future storage will use dict format." + "Retrieved code_verifier in legacy plain-string format. Future storage will use dict format." ) else: # Defer the detailed ERROR log to the strict-mode branch below @@ -3621,12 +3363,8 @@ class SSOAuthenticationHandler: In strict mode (PKCE_STRICT_CACHE_MISS=true) raises ProxyException. Otherwise logs a warning and returns (token exchange proceeds without verifier). """ - active_cache = ( - redis_usage_cache if redis_usage_cache is not None else user_api_key_cache - ) - strict_cache_miss = ( - os.getenv("PKCE_STRICT_CACHE_MISS", "false").lower() == "true" - ) + active_cache = redis_usage_cache if redis_usage_cache is not None else user_api_key_cache + strict_cache_miss = os.getenv("PKCE_STRICT_CACHE_MISS", "false").lower() == "true" if strict_cache_miss: if empty_value_in_dict: await SSOAuthenticationHandler._delete_pkce_verifier(cache_key) @@ -3670,8 +3408,7 @@ class SSOAuthenticationHandler: "Configure Redis so all proxy instances share the PKCE verifier." ) verbose_proxy_logger.error( - "PKCE is enabled but no verifier found in cache for state '%s'. " - "%s Cache type: %s.", + "PKCE is enabled but no verifier found in cache for state '%s'. %s Cache type: %s.", state, cause, type(active_cache).__name__, @@ -3729,17 +3466,11 @@ class SSOAuthenticationHandler: """ # Generate a cryptographically random code_verifier (43 characters) # Using 32 random bytes which becomes 43 characters when base64-url-encoded - code_verifier = ( - base64.urlsafe_b64encode(secrets.token_bytes(32)) - .decode("utf-8") - .rstrip("=") - ) + code_verifier = base64.urlsafe_b64encode(secrets.token_bytes(32)).decode("utf-8").rstrip("=") # Generate code_challenge using S256 method (SHA256) code_challenge_bytes = hashlib.sha256(code_verifier.encode("utf-8")).digest() - code_challenge = ( - base64.urlsafe_b64encode(code_challenge_bytes).decode("utf-8").rstrip("=") - ) + code_challenge = base64.urlsafe_b64encode(code_challenge_bytes).decode("utf-8").rstrip("=") return code_verifier, code_challenge @@ -3794,9 +3525,7 @@ class SSOAuthenticationHandler: "token endpoint returned HTTP 200 but no access_token " f"(response keys: {sorted(token_response.keys())})" ) - verbose_proxy_logger.error( - "Token response missing or null access_token. detail=%s", detail - ) + verbose_proxy_logger.error("Token response missing or null access_token. detail=%s", detail) raise ProxyException( message=f"Token exchange failed: {detail}", type=ProxyErrorTypes.auth_error, @@ -3851,9 +3580,7 @@ class SSOAuthenticationHandler: if not include_client_id: # Use Basic Auth only when a secret is available; public PKCE clients omit it. if client_secret: - credentials = base64.b64encode( - f"{client_id}:{client_secret}".encode() - ).decode() + credentials = base64.b64encode(f"{client_id}:{client_secret}".encode()).decode() request_headers["Authorization"] = f"Basic {credentials}" else: token_data["client_id"] = client_id @@ -3862,9 +3589,7 @@ class SSOAuthenticationHandler: if client_secret: token_data["client_secret"] = client_secret - http_client = get_async_httpx_client( - llm_provider=httpxSpecialProvider.SSO_HANDLER - ) + http_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.SSO_HANDLER) try: response = await http_client.post( url=token_endpoint, @@ -3954,9 +3679,7 @@ class SSOAuthenticationHandler: if userinfo_endpoint: try: - client = get_async_httpx_client( - llm_provider=httpxSpecialProvider.SSO_HANDLER - ) + client = get_async_httpx_client(llm_provider=httpxSpecialProvider.SSO_HANDLER) resp = await client.get( url=userinfo_endpoint, headers={ @@ -3991,9 +3714,7 @@ class SSOAuthenticationHandler: resp.text[:500], ) except Exception as e: - verbose_proxy_logger.warning( - "Userinfo endpoint error: %s, falling back to id_token", e - ) + verbose_proxy_logger.warning("Userinfo endpoint error: %s, falling back to id_token", e) # Only fall back to id_token when the userinfo request failed (None). # Empty dict ({}) and JSON null are both treated as failure (set to None above) since @@ -4007,9 +3728,7 @@ class SSOAuthenticationHandler: # jwt.decode returned an empty dict (payload-free JWT or provider bug). # Treat this the same as a missing userinfo — the session would have no # identity claims, which is equivalent to a broken session. - verbose_proxy_logger.warning( - "id_token decoded to an empty payload — treating as failure." - ) + verbose_proxy_logger.warning("id_token decoded to an empty payload — treating as failure.") userinfo = None except Exception as decode_err: verbose_proxy_logger.error("Failed to decode id_token: %s", decode_err) @@ -4037,7 +3756,9 @@ class SSOAuthenticationHandler: "and id_token decoded to an empty payload — no identity claims available" ) else: - detail = "no userinfo endpoint is configured (GENERIC_USERINFO_ENDPOINT) and no id_token was present" + detail = ( + "no userinfo endpoint is configured (GENERIC_USERINFO_ENDPOINT) and no id_token was present" + ) raise ProxyException( message=f"SSO user info unavailable: {detail}.", type=ProxyErrorTypes.auth_error, @@ -4113,9 +3834,7 @@ class MicrosoftSSOHandler: ) # Extract app roles from the id_token JWT - app_roles = MicrosoftSSOHandler.get_app_roles_from_id_token( - id_token=microsoft_sso.id_token - ) + app_roles = MicrosoftSSOHandler.get_app_roles_from_id_token(id_token=microsoft_sso.id_token) verbose_proxy_logger.debug(f"Extracted app roles from id_token: {app_roles}") # Combine groups and app roles @@ -4126,20 +3845,14 @@ class MicrosoftSSOHandler: role = get_litellm_user_role(role_str) if role is not None: user_role = role - verbose_proxy_logger.debug( - f"Found valid LitellmUserRoles '{role.value}' in app_roles" - ) + verbose_proxy_logger.debug(f"Found valid LitellmUserRoles '{role.value}' in app_roles") break - verbose_proxy_logger.debug( - f"Combined team_ids (groups + app roles): {user_team_ids}" - ) + verbose_proxy_logger.debug(f"Combined team_ids (groups + app roles): {user_team_ids}") # if user is trying to get the raw sso response for debugging, return the raw sso response if return_raw_sso_response: - original_msft_result[MicrosoftSSOHandler.GRAPH_API_RESPONSE_KEY] = ( - user_team_ids - ) + original_msft_result[MicrosoftSSOHandler.GRAPH_API_RESPONSE_KEY] = user_team_ids original_msft_result["app_roles"] = app_roles return original_msft_result or {} @@ -4159,9 +3872,7 @@ class MicrosoftSSOHandler: response = response or {} verbose_proxy_logger.debug(f"Microsoft SSO Callback Response: {response}") openid_response = CustomOpenID( - email=normalize_email( - response.get(MICROSOFT_USER_EMAIL_ATTRIBUTE) or response.get("mail") - ), + email=normalize_email(response.get(MICROSOFT_USER_EMAIL_ATTRIBUTE) or response.get("mail")), display_name=response.get(MICROSOFT_USER_DISPLAY_NAME_ATTRIBUTE), provider="microsoft", id=response.get(MICROSOFT_USER_ID_ATTRIBUTE), @@ -4203,14 +3914,10 @@ class MicrosoftSSOHandler: roles = decoded_token.get("app_roles", []) or decoded_token.get("roles", []) if roles and isinstance(roles, list): - verbose_proxy_logger.debug( - f"Found {len(roles)} app role(s) in id_token: {roles}" - ) + verbose_proxy_logger.debug(f"Found {len(roles)} app role(s) in id_token: {roles}") return roles else: - verbose_proxy_logger.debug( - "No app roles found in id_token or roles claim is not a list" - ) + verbose_proxy_logger.debug("No app roles found in id_token or roles claim is not a list") return [] except Exception as e: @@ -4231,9 +3938,7 @@ class MicrosoftSSOHandler: List[str]: List of group IDs the user belongs to """ try: - async_client = get_async_httpx_client( - llm_provider=httpxSpecialProvider.SSO_HANDLER - ) + async_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.SSO_HANDLER) # Handle MSFT Enterprise Application Groups service_principal_id = os.getenv("MICROSOFT_SERVICE_PRINCIPAL_ID", None) @@ -4248,9 +3953,7 @@ class MicrosoftSSOHandler: async_client=async_client, access_token=access_token, ) - verbose_proxy_logger.debug( - f"Service principal group IDs: {service_principal_group_ids}" - ) + verbose_proxy_logger.debug(f"Service principal group IDs: {service_principal_group_ids}") if len(service_principal_group_ids) > 0: await MicrosoftSSOHandler.create_litellm_teams_from_service_principal_team_ids( service_principal_teams=service_principal_teams, @@ -4258,44 +3961,30 @@ class MicrosoftSSOHandler: # Fetch user membership from Microsoft Graph API all_group_ids = [] - next_link: Optional[str] = ( - MicrosoftSSOHandler.graph_api_user_groups_endpoint - ) + next_link: Optional[str] = MicrosoftSSOHandler.graph_api_user_groups_endpoint auth_headers = {"Authorization": f"Bearer {access_token}"} page_count = 0 - while ( - next_link is not None - and page_count < MicrosoftSSOHandler.MAX_GRAPH_API_PAGES - ): + while next_link is not None and page_count < MicrosoftSSOHandler.MAX_GRAPH_API_PAGES: group_ids, next_link = await MicrosoftSSOHandler.fetch_and_parse_groups( url=next_link, headers=auth_headers, async_client=async_client ) all_group_ids.extend(group_ids) page_count += 1 - if ( - next_link is not None - and page_count >= MicrosoftSSOHandler.MAX_GRAPH_API_PAGES - ): + if next_link is not None and page_count >= MicrosoftSSOHandler.MAX_GRAPH_API_PAGES: verbose_proxy_logger.warning( f"Reached maximum page limit of {MicrosoftSSOHandler.MAX_GRAPH_API_PAGES}. Some groups may not be included." ) # If service_principal_group_ids is not empty, only return group_ids that are in both all_group_ids and service_principal_group_ids if service_principal_group_ids and len(service_principal_group_ids) > 0: - all_group_ids = [ - group_id - for group_id in all_group_ids - if group_id in service_principal_group_ids - ] + all_group_ids = [group_id for group_id in all_group_ids if group_id in service_principal_group_ids] return all_group_ids except Exception as e: - verbose_proxy_logger.error( - f"Error getting user groups from Microsoft Graph API: {e}" - ) + verbose_proxy_logger.error(f"Error getting user groups from Microsoft Graph API: {e}") return [] @staticmethod @@ -4305,12 +3994,8 @@ class MicrosoftSSOHandler: """Helper function to fetch and parse group data from a URL""" response = await async_client.get(url, headers=headers) response_json = response.json() - response_typed = await MicrosoftSSOHandler._cast_graph_api_response_dict( - response=response_json - ) - group_ids = MicrosoftSSOHandler._get_group_ids_from_graph_api_response( - response=response_typed - ) + response_typed = await MicrosoftSSOHandler._cast_graph_api_response_dict(response=response_json) + group_ids = MicrosoftSSOHandler._get_group_ids_from_graph_api_response(response=response_typed) return group_ids, response_typed.get("odata_nextLink") @staticmethod @@ -4371,9 +4056,7 @@ class MicrosoftSSOHandler: response = await async_client.get(url, headers=headers) response_json = response.json() - verbose_proxy_logger.debug( - f"Response from service principal app role assigned to: {response_json}" - ) + verbose_proxy_logger.debug(f"Response from service principal app role assigned to: {response_json}") group_ids: List[str] = [] service_principal_teams: List[MicrosoftServicePrincipalTeam] = [] @@ -4400,14 +4083,10 @@ class MicrosoftSSOHandler: When a user sets a `SERVICE_PRINCIPAL_ID` in the env, litellm will fetch groups under that service principal and create Litellm Teams from them """ - verbose_proxy_logger.debug( - f"Creating Litellm Teams from Service Principal Teams: {service_principal_teams}" - ) + verbose_proxy_logger.debug(f"Creating Litellm Teams from Service Principal Teams: {service_principal_teams}") for service_principal_team in service_principal_teams: litellm_team_id: Optional[str] = service_principal_team.get("principalId") - litellm_team_name: Optional[str] = service_principal_team.get( - "principalDisplayName" - ) + litellm_team_name: Optional[str] = service_principal_team.get("principalDisplayName") if not litellm_team_id: verbose_proxy_logger.debug( f"Skipping team creation for {litellm_team_name} because it has no principalId" @@ -4482,11 +4161,7 @@ async def debug_sso_login(request: Request): generic_client_id = os.getenv("GENERIC_CLIENT_ID", None) ####### Check if user is a Enterprise / Premium User ####### - if ( - microsoft_client_id is not None - or google_client_id is not None - or generic_client_id is not None - ): + if microsoft_client_id is not None or google_client_id is not None or generic_client_id is not None: if premium_user is not True: raise ProxyException( message="You must be a LiteLLM Enterprise user to use SSO. If you have a license please set `LITELLM_LICENSE` in your env. If you want to obtain a license meet with us here: https://enterprise.litellm.ai/demo You are seeing this error message because You set one of `MICROSOFT_CLIENT_ID`, `GOOGLE_CLIENT_ID`, or `GENERIC_CLIENT_ID` in your env. Please unset this", @@ -4545,9 +4220,7 @@ async def debug_sso_callback(request: Request): prisma_client=prisma_client, user_api_key_cache=user_api_key_cache, litellm_jwtauth=LiteLLM_JWTAuth( - team_ids_jwt_field=general_settings.get("ui_access_mode", {}).get( - "sso_group_jwt_field", None - ), + team_ids_jwt_field=general_settings.get("ui_access_mode", {}).get("sso_group_jwt_field", None), ), leeway=0, ) @@ -4581,14 +4254,12 @@ async def debug_sso_callback(request: Request): ) elif generic_client_id is not None: - result, received_response, access_token_payload = ( - await get_generic_sso_response( - request=request, - jwt_handler=jwt_handler, - generic_client_id=generic_client_id, - redirect_url=redirect_url, - sso_jwt_handler=sso_jwt_handler, - ) + result, received_response, access_token_payload = await get_generic_sso_response( + request=request, + jwt_handler=jwt_handler, + generic_client_id=generic_client_id, + redirect_url=redirect_url, + sso_jwt_handler=sso_jwt_handler, ) # If result is None, return a basic error message @@ -4619,16 +4290,8 @@ async def debug_sso_callback(request: Request): # Defense-in-depth: ensure no bearer tokens leak into the rendered HTML even if # a non-conforming IdP places them in its userinfo response. - safe_raw_claims = { - k: v - for k, v in (received_response or {}).items() - if k not in _OAUTH_TOKEN_FIELDS - } - safe_access_token_claims = { - k: v - for k, v in (access_token_payload or {}).items() - if k not in _OAUTH_TOKEN_FIELDS - } + safe_raw_claims = {k: v for k, v in (received_response or {}).items() if k not in _OAUTH_TOKEN_FIELDS} + safe_access_token_claims = {k: v for k, v in (access_token_payload or {}).items() if k not in _OAUTH_TOKEN_FIELDS} sso_payload = { "parsed_by_proxy": filtered_result, @@ -4637,9 +4300,7 @@ async def debug_sso_callback(request: Request): } # Replace the placeholder in the template with the actual data - sso_payload_json = json.dumps(sso_payload, indent=2, default=str).replace( - " None: + def normalize_authorized_vector_store_id(cls, vector_store_opts: dict[str, object]) -> None: """ Rewrite the vector_store config so `vector_store_id` matches the actual write target before the proxy authorizes it. diff --git a/litellm/rag/ingestion/milvus_ingestion.py b/litellm/rag/ingestion/milvus_ingestion.py index aa21133816e..e5079022168 100644 --- a/litellm/rag/ingestion/milvus_ingestion.py +++ b/litellm/rag/ingestion/milvus_ingestion.py @@ -76,9 +76,9 @@ class MilvusRAGIngestion(BaseRAGIngestion): if not self.embedding_config: self.embedding_config = {"model": "text-embedding-3-small"} - self.collection_name = self.vector_store_config.get( - "collection_name" - ) or self.vector_store_config.get("vector_store_id") + self.collection_name = self.vector_store_config.get("collection_name") or self.vector_store_config.get( + "vector_store_id" + ) if not self.collection_name: raise ValueError( "Milvus RAG ingestion requires 'collection_name' (or 'vector_store_id') in the vector_store config." @@ -94,34 +94,18 @@ class MilvusRAGIngestion(BaseRAGIngestion): self.api_base = self.api_base.rstrip("/") config_api_key = self.vector_store_config.get("api_key") - self.api_key = ( - config_api_key - if config_api_base - else config_api_key or get_secret_str("MILVUS_API_KEY") - ) - self.vector_field = self.vector_store_config.get( - "vector_field", MILVUS_DEFAULT_VECTOR_FIELD - ) - self.text_field = self.vector_store_config.get( - "text_field", MILVUS_DEFAULT_TEXT_FIELD - ) - self.metric_type = self.vector_store_config.get( - "metric_type", MILVUS_DEFAULT_METRIC_TYPE - ) + self.api_key = config_api_key if config_api_base else config_api_key or get_secret_str("MILVUS_API_KEY") + self.vector_field = self.vector_store_config.get("vector_field", MILVUS_DEFAULT_VECTOR_FIELD) + self.text_field = self.vector_store_config.get("text_field", MILVUS_DEFAULT_TEXT_FIELD) + self.metric_type = self.vector_store_config.get("metric_type", MILVUS_DEFAULT_METRIC_TYPE) self.db_name = get_secret_str("MILVUS_DB_NAME") self.partition_name = get_secret_str("MILVUS_PARTITION_NAME") - self.auto_create_collection = self.vector_store_config.get( - "auto_create_collection", True - ) + self.auto_create_collection = self.vector_store_config.get("auto_create_collection", True) - self.async_httpx_client = get_async_httpx_client( - llm_provider=httpxSpecialProvider.RAG - ) + self.async_httpx_client = get_async_httpx_client(llm_provider=httpxSpecialProvider.RAG) @classmethod - def normalize_authorized_vector_store_id( - cls, vector_store_opts: dict[str, object] - ) -> None: + def normalize_authorized_vector_store_id(cls, vector_store_opts: dict[str, object]) -> None: """ Milvus resolves its write target from `collection_name` first (falling back to `vector_store_id`). Always mirror `collection_name` onto @@ -167,9 +151,7 @@ class MilvusRAGIngestion(BaseRAGIngestion): async def _post(self, path: str, body: dict[str, object]) -> dict[str, object]: url = f"{self.api_base}{path}" - response = await self.async_httpx_client.post( - url, json=body, headers=self._headers() - ) + response = await self.async_httpx_client.post(url, json=body, headers=self._headers()) response.raise_for_status() data = response.json() # Milvus REST returns {"code": 0, "data": ...} on success, @@ -261,8 +243,6 @@ class MilvusRAGIngestion(BaseRAGIngestion): body["partitionName"] = self.partition_name await self._post("/v2/vectordb/entities/insert", body) - verbose_logger.info( - f"Inserted {len(rows)} vectors into Milvus collection '{self.collection_name}'" - ) + verbose_logger.info(f"Inserted {len(rows)} vectors into Milvus collection '{self.collection_name}'") return self.collection_name, filename diff --git a/litellm/router_utils/cooldown_callbacks.py b/litellm/router_utils/cooldown_callbacks.py index bccb25cd407..068790fbef2 100644 --- a/litellm/router_utils/cooldown_callbacks.py +++ b/litellm/router_utils/cooldown_callbacks.py @@ -55,12 +55,7 @@ async def router_cooldown_event_callback( except Exception: pass - _api_base = ( - litellm.get_api_base( - model=litellm_model_name, optional_params=temp_litellm_params - ) - or "" - ) + _api_base = litellm.get_api_base(model=litellm_model_name, optional_params=temp_litellm_params) or "" # get the prometheus logger from in memory loggers prometheusLogger: Optional[PrometheusLogger] = _get_prometheus_logger_from_callbacks() diff --git a/litellm/types/guardrails.py b/litellm/types/guardrails.py index 1b0140e6b8c..1792683af60 100644 --- a/litellm/types/guardrails.py +++ b/litellm/types/guardrails.py @@ -126,6 +126,7 @@ class SupportedGuardrailIntegrations(Enum): VIGIL_GUARD = "vigil_guard" REPELLOAI = "repelloai" BIAS_HALLUCINATION_ESTIMATOR = "bias_hallucination_estimator" + HEADROOM = "headroom" class Role(Enum): @@ -387,12 +388,8 @@ class PresidioConfigModel(PresidioPresidioConfigModelUserInterface): mock_redacted_text: Optional[dict] = Field(default=None, description="Mock redacted text for testing") -BedrockChecksContentFilterCategory = Literal[ - "VIOLENCE", "HATE", "SEXUAL", "MISCONDUCT", "INSULTS" -] -BedrockChecksPromptAttackCategory = Literal[ - "JAILBREAK", "PROMPT_INJECTION", "PROMPT_LEAKAGE" -] +BedrockChecksContentFilterCategory = Literal["VIOLENCE", "HATE", "SEXUAL", "MISCONDUCT", "INSULTS"] +BedrockChecksPromptAttackCategory = Literal["JAILBREAK", "PROMPT_INJECTION", "PROMPT_LEAKAGE"] BedrockChecksSensitiveInformationEntity = Literal[ "ADDRESS", "AGE", @@ -464,14 +461,9 @@ class BedrockChecksConfigModel(BaseModel): @model_validator(mode="after") def _require_at_least_one_check(self) -> "BedrockChecksConfigModel": - if ( - self.contentFilter is None - and self.promptAttack is None - and self.sensitiveInformation is None - ): + if self.contentFilter is None and self.promptAttack is None and self.sensitiveInformation is None: raise ValueError( - "Bedrock 'checks' must enable at least one of: contentFilter, " - "promptAttack, sensitiveInformation." + "Bedrock 'checks' must enable at least one of: contentFilter, promptAttack, sensitiveInformation." ) return self @@ -498,12 +490,8 @@ class BedrockGuardrailConfigModel(BaseModel): aws_web_identity_token: Optional[str] = Field( default=None, description="Web identity token for AWS role assumption" ) - aws_sts_endpoint: Optional[str] = Field( - default=None, description="AWS STS endpoint URL" - ) - aws_bedrock_runtime_endpoint: Optional[str] = Field( - default=None, description="AWS Bedrock runtime endpoint URL" - ) + aws_sts_endpoint: Optional[str] = Field(default=None, description="AWS STS endpoint URL") + aws_bedrock_runtime_endpoint: Optional[str] = Field(default=None, description="AWS Bedrock runtime endpoint URL") checks: BedrockChecksConfigModel | None = Field( default=None, description="Inline safeguards for the resource-less InvokeGuardrailChecks API " diff --git a/litellm/types/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator.py b/litellm/types/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator.py index 6fc2ccf58b0..6da68f50afc 100644 --- a/litellm/types/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator.py +++ b/litellm/types/proxy/guardrails/guardrail_hooks/bias_hallucination_estimator.py @@ -8,9 +8,7 @@ from litellm.types.proxy.guardrails.guardrail_hooks.base import GuardrailConfigM class BiasAnalysis(BaseModel): """Model representing the result of a bias analysis.""" - bias_detected: bool = Field( - default=False, description="Indicates if bias was detected in the text" - ) + bias_detected: bool = Field(default=False, description="Indicates if bias was detected in the text") score: float = Field( default=0.0, ge=0.0, @@ -82,9 +80,7 @@ class UncertaintyAnalysis(BaseModel): class RiskScore(BaseModel): """Model representing the overall risk score combining bias and hallucination.""" - overall_risk_percentage: int = Field( - default=0, ge=0, le=100, description="Overall risk percentage (0-100)" - ) + overall_risk_percentage: int = Field(default=0, ge=0, le=100, description="Overall risk percentage (0-100)") bias_score: float = Field(default=0.0, ge=0.0, le=1.0) hallucination_score: float = Field(default=0.0, ge=0.0, le=1.0) uncertainty_score: float = Field(default=0.0, ge=0.0, le=1.0) @@ -92,9 +88,7 @@ class RiskScore(BaseModel): recommendation: Literal["pass", "flag", "block"] = Field(default="pass") -class BiasHallucinationEstimatorConfigModel( - GuardrailConfigModel -): # pyright: ignore[reportMissingTypeArgument] +class BiasHallucinationEstimatorConfigModel(GuardrailConfigModel): # pyright: ignore[reportMissingTypeArgument] """Configuration schema for the native bias and hallucination estimator.""" bias_threshold: float = Field(default=0.5, ge=0.0, le=1.0) diff --git a/litellm/types/rag.py b/litellm/types/rag.py index 4793fa32b13..40df4be8f7a 100644 --- a/litellm/types/rag.py +++ b/litellm/types/rag.py @@ -203,9 +203,7 @@ class MilvusVectorStoreOptions(TypedDict, total=False): vector_field: Optional[str] # Embedding field name (default: "vector") text_field: Optional[str] # Chunk text field name (default: "text") metric_type: Optional[str] # Distance metric (default: "COSINE") - auto_create_collection: Optional[ - bool - ] # Create collection if missing (default: True) + auto_create_collection: Optional[bool] # Create collection if missing (default: True) # Union type for vector store options