mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
f0f0cb71f0
commit
a09c11fcaf
23 changed files with 338 additions and 924 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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 []
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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 "
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue