fix: rebase on litellm_internal_staging, fix ruff format and SSO budget tests

- Resolve merge conflicts from rebasing onto litellm_internal_staging
- Add HEADROOM to SupportedGuardrailIntegrations enum (dropped during conflict resolution)
- Fix cli_poll_key to pass max_budget to get_cli_jwt_auth_token: look up user/team budgets and fall back to max_ui_session_budget when neither has one
- Run ruff format on all changed files to fix lint failures

Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
This commit is contained in:
Sameer Kankute 2026-06-29 09:28:30 +05:30
parent f0f0cb71f0
commit a09c11fcaf
No known key found for this signature in database
23 changed files with 338 additions and 924 deletions

View file

@ -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

View file

@ -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()

View file

@ -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,

View file

@ -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
)

View file

@ -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

View file

@ -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)

View file

@ -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(

View file

@ -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(

View file

@ -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

View file

@ -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

View file

@ -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 []

View file

@ -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

View file

@ -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)

View file

@ -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,
)

View file

@ -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:

View file

@ -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)

File diff suppressed because it is too large Load diff

View file

@ -102,9 +102,7 @@ class BaseRAGIngestion(ABC):
return self.vector_store_config.get("custom_llm_provider", "openai")
@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:
"""
Rewrite the vector_store config so `vector_store_id` matches the actual
write target before the proxy authorizes it.

View file

@ -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

View file

@ -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()

View file

@ -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 "

View file

@ -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)

View file

@ -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