diff --git a/.github/workflows/docker.yml b/.github/workflows/docker.yml index 7b29b7d..cba7c08 100644 --- a/.github/workflows/docker.yml +++ b/.github/workflows/docker.yml @@ -6,7 +6,7 @@ on: workflow_dispatch: inputs: version: - description: "Version to publish (e.g. 0.0.14). Must match pyproject.toml." + description: "Version to publish (e.g. 0.0.15). Must match pyproject.toml." required: true type: string branch: diff --git a/README.md b/README.md index d78fd35..3e50a02 100644 --- a/README.md +++ b/README.md @@ -204,7 +204,8 @@ Open `http://localhost:8080/admin` to see the gateway dashboard. It includes: - Login with `admin` and `viewer` roles. - Per-client request, token, cost, model, provider, and cache-hit usage. - Provider health and routing status. -- Admin-only client token creation, allowed-model updates, enable/disable, and rotation. +- Admin-only client token creation, allowed-model updates, per-client policy actions, enable/disable, and rotation. +- Admin-only upstream provider key updates with masked display and encrypted SQLite storage. - One-time token reveal on create or rotate; SentinelGuard stores only a hash. For local testing, the fallback credentials are `admin` / `sentinelguard` and @@ -215,6 +216,13 @@ client tokens are separate from upstream LLM provider keys. Apps, SDKs, IDEs, EKS services, or EC2 services use the generated `sgw_...` token as their API key when their base URL points to SentinelGuard. +Set a stable `SENTINELGUARD_ENCRYPTION_KEY` if you want the dashboard to store +upstream provider API keys. Generate one with +`sentinelguard token --prefix sgencrypt`, or let +`sentinelguard init --with-env` write it into the local `.env` file. Without it, +provider keys should come from env vars, Docker secrets, Kubernetes Secrets, or +your secret manager. + ## Connect Apps, Services, And IDEs Point clients to SentinelGuard instead of directly to the model provider: diff --git a/docker-compose.yml b/docker-compose.yml index ea281e7..080ab2f 100644 --- a/docker-compose.yml +++ b/docker-compose.yml @@ -19,6 +19,7 @@ services: HUGGINGFACE_API_KEY: ${HUGGINGFACE_API_KEY:-} SENTINELGUARD_GATEWAY_API_KEY: ${SENTINELGUARD_GATEWAY_API_KEY:-} SENTINELGUARD_AUDIT_SALT: ${SENTINELGUARD_AUDIT_SALT:-local-dev-salt} + SENTINELGUARD_ENCRYPTION_KEY: ${SENTINELGUARD_ENCRYPTION_KEY:-} SENTINELGUARD_ADMIN_USERNAME: ${SENTINELGUARD_ADMIN_USERNAME:-admin} SENTINELGUARD_ADMIN_PASSWORD: ${SENTINELGUARD_ADMIN_PASSWORD:-sentinelguard} SENTINELGUARD_VIEWER_USERNAME: ${SENTINELGUARD_VIEWER_USERNAME:-viewer} diff --git a/docs/container-images.md b/docs/container-images.md index 8a18375..f07eb7b 100644 --- a/docs/container-images.md +++ b/docs/container-images.md @@ -25,7 +25,7 @@ Use `latest` only for local testing or demos. docker run --rm -p 8080:8080 \ -e OPENAI_API_KEY="$OPENAI_API_KEY" \ -e SENTINELGUARD_GATEWAY_API_KEY="$SENTINELGUARD_GATEWAY_API_KEY" \ - aitechnav/sentinelguard:0.0.14 + aitechnav/sentinelguard:0.0.15 ``` Then point OpenAI-compatible clients to: diff --git a/docs/deployment.md b/docs/deployment.md index 954faa6..dfa9a32 100644 --- a/docs/deployment.md +++ b/docs/deployment.md @@ -35,6 +35,13 @@ Base URL: http://sentinelguard-gateway:8080/v1 API key: the same sgw_... value from SENTINELGUARD_GATEWAY_API_KEY ``` +For Docker Compose, put upstream provider keys such as `OPENAI_API_KEY` in the +Compose `.env` file or Docker secrets so only the SentinelGuard container +receives them. To update provider keys from the dashboard instead, also set a +stable `SENTINELGUARD_ENCRYPTION_KEY`. `sentinelguard init --with-env` generates +one for local Compose. Dashboard-entered provider keys are then stored encrypted +in the gateway SQLite volume and shown only as masked hints. + ## Kubernetes ```bash @@ -50,6 +57,12 @@ Base URL: http://sentinelguard-gateway.sentinelguard.svc.cluster.local:8080/v1 API key: the same sgw_... value from SENTINELGUARD_GATEWAY_API_KEY ``` +For Kubernetes, keep upstream provider keys in Kubernetes Secrets, external +secret operators, or your cloud secret manager and expose them to SentinelGuard +as environment variables referenced by `api_key_env`. Dashboard-entered provider +keys are supported for self-managed deployments, but they are stored in the +gateway SQLite state volume, not written back to a Kubernetes Secret. + For EC2, ECS, another EKS cluster, or another VPC, expose the gateway through a private DNS name, internal load balancer, PrivateLink, VPN, or peering route: diff --git a/docs/gateway.md b/docs/gateway.md index bef41e1..1e66c82 100644 --- a/docs/gateway.md +++ b/docs/gateway.md @@ -141,6 +141,30 @@ client in `/admin`, edit **Allowed models**, and save. For example: `sentinel-auto, fast-chat, smart-chat, private-chat`. This updates the existing client token policy; it does not require rotating the token. +Admins can also update policy actions per client from the same dashboard form. +Use this when one app should block attacks and secrets but redact PII, while +another app should audit PII or block PCI data. Supported actions are `block`, +`redact`, `audit`, and `allow` for attack, secret, PII, PCI, PHI, and other +scanner categories. + +Admins can update upstream provider API keys from the dashboard when encrypted +provider-secret storage is enabled. Set one stable encryption key on the +gateway process: + +```bash +export SENTINELGUARD_ENCRYPTION_KEY="$(sentinelguard token --prefix sgencrypt)" +``` + +`sentinelguard init --with-env` also generates this value in the local `.env` +file used by Docker Compose. + +Dashboard-entered OpenAI, Anthropic, Gemini, or other provider keys are stored +encrypted in the gateway SQLite database and shown only as a masked hint. They +override the environment/YAML key for that provider route until the dashboard +key is removed. SentinelGuard cannot rotate provider API keys because those are +issued by the provider; it can update, replace, test, or remove the configured +key. + A dashboard-managed client can also rotate its own token while it still has a valid current token: diff --git a/pyproject.toml b/pyproject.toml index 845f411..5f3d174 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -4,7 +4,7 @@ build-backend = "setuptools.build_meta" [project] name = "sentinelguard" -version = "0.0.14" +version = "0.0.15" description = "A comprehensive, production-ready LLM security and guardrails framework" readme = "README.md" license = "Apache-2.0" @@ -56,6 +56,7 @@ gateway = [ "fastapi>=0.100.0", "uvicorn>=0.23.0", "httpx>=0.24.0", + "cryptography>=41.0.0", "redis>=5.0.0", "websockets>=12.0", ] diff --git a/requirements.txt b/requirements.txt index 065d609..9c2c172 100644 --- a/requirements.txt +++ b/requirements.txt @@ -21,6 +21,7 @@ regex>=2023.0 # fastapi>=0.100.0 # uvicorn>=0.23.0 # httpx>=0.24.0 +# cryptography>=41.0.0 # encrypted dashboard-managed provider secrets # Monitoring (optional) # opentelemetry-api>=1.20.0 diff --git a/sentinelguard/__init__.py b/sentinelguard/__init__.py index 009c878..df68ecd 100644 --- a/sentinelguard/__init__.py +++ b/sentinelguard/__init__.py @@ -36,7 +36,7 @@ guard = SentinelGuard(config=config) """ -__version__ = "0.0.14" +__version__ = "0.0.15" __author__ = "SentinelGuard Contributors" from sentinelguard.core.guard import SentinelGuard diff --git a/sentinelguard/cli/bootstrap.py b/sentinelguard/cli/bootstrap.py index fa38bc8..2e52c8f 100644 --- a/sentinelguard/cli/bootstrap.py +++ b/sentinelguard/cli/bootstrap.py @@ -235,6 +235,7 @@ def _env_example() -> str: # Generate local SentinelGuard tokens and dashboard passwords with: # sentinelguard token # sentinelguard token --prefix sgaudit + # sentinelguard token --prefix sgencrypt # sentinelguard token --prefix sgadmin # sentinelguard token --prefix sgviewer @@ -243,6 +244,7 @@ def _env_example() -> str: SENTINELGUARD_IMAGE=sentinelguard-gateway:local SENTINELGUARD_GATEWAY_API_KEY= SENTINELGUARD_AUDIT_SALT= + SENTINELGUARD_ENCRYPTION_KEY= SENTINELGUARD_ADMIN_USERNAME=admin SENTINELGUARD_ADMIN_PASSWORD= SENTINELGUARD_VIEWER_USERNAME=viewer @@ -268,6 +270,7 @@ def _env_example() -> str: def _env_file() -> str: gateway_token = generate_gateway_token() audit_salt = generate_gateway_token(prefix="sgaudit") + encryption_key = generate_gateway_token(prefix="sgencrypt") admin_password = generate_gateway_token(prefix="sgadmin") viewer_password = generate_gateway_token(prefix="sgviewer") return dedent( @@ -280,6 +283,7 @@ def _env_file() -> str: SENTINELGUARD_IMAGE=sentinelguard-gateway:local SENTINELGUARD_GATEWAY_API_KEY={gateway_token} SENTINELGUARD_AUDIT_SALT={audit_salt} + SENTINELGUARD_ENCRYPTION_KEY={encryption_key} SENTINELGUARD_ADMIN_USERNAME=admin SENTINELGUARD_ADMIN_PASSWORD={admin_password} SENTINELGUARD_VIEWER_USERNAME=viewer @@ -359,6 +363,7 @@ def _docker_compose(docker_image: str) -> str: OLLAMA_API_KEY: ${{OLLAMA_API_KEY:-}} SENTINELGUARD_GATEWAY_API_KEY: ${{SENTINELGUARD_GATEWAY_API_KEY:-}} SENTINELGUARD_AUDIT_SALT: ${{SENTINELGUARD_AUDIT_SALT:-local-dev-audit-salt}} + SENTINELGUARD_ENCRYPTION_KEY: ${{SENTINELGUARD_ENCRYPTION_KEY:-}} SENTINELGUARD_ADMIN_USERNAME: ${{SENTINELGUARD_ADMIN_USERNAME:-admin}} SENTINELGUARD_ADMIN_PASSWORD: ${{SENTINELGUARD_ADMIN_PASSWORD:-sentinelguard}} SENTINELGUARD_VIEWER_USERNAME: ${{SENTINELGUARD_VIEWER_USERNAME:-viewer}} @@ -425,6 +430,7 @@ def _readme(profile: str) -> str: python -m pip install "sentinelguard[gateway,monitoring]" export SENTINELGUARD_GATEWAY_API_KEY="$(sentinelguard token)" export SENTINELGUARD_AUDIT_SALT="$(sentinelguard token --prefix sgaudit)" + export SENTINELGUARD_ENCRYPTION_KEY="$(sentinelguard token --prefix sgencrypt)" export OPENAI_API_KEY="your-provider-key" sentinelguard gateway \\ --config sentinelguard.yaml \\ diff --git a/sentinelguard/gateway/admin.py b/sentinelguard/gateway/admin.py index e2764b9..8d4ffce 100644 --- a/sentinelguard/gateway/admin.py +++ b/sentinelguard/gateway/admin.py @@ -11,11 +11,11 @@ import sqlite3 import threading import time -from dataclasses import dataclass +from dataclasses import dataclass, field from pathlib import Path from typing import Any, Optional -from sentinelguard.gateway.config import GatewayConfig +from sentinelguard.gateway.config import GatewayConfig, normalize_policy_actions from sentinelguard.gateway.operations import GatewayClient ADMIN_COOKIE_NAME = "sentinelguard_admin_session" @@ -45,6 +45,7 @@ class StoredGatewayToken: team_id: Optional[str] = None user_id: Optional[str] = None allowed_models: tuple[str, ...] = () + policy_actions: dict[str, str] = field(default_factory=dict) max_requests: Optional[int] = None max_tokens: Optional[int] = None max_budget: Optional[float] = None @@ -65,6 +66,7 @@ def to_dict(self) -> dict[str, Any]: "team_id": self.team_id, "user_id": self.user_id, "allowed_models": list(self.allowed_models), + "policy_actions": normalize_policy_actions(self.policy_actions), "max_requests": self.max_requests, "max_tokens": self.max_tokens, "max_budget": self.max_budget, @@ -78,6 +80,26 @@ def to_dict(self) -> dict[str, Any]: } +@dataclass(frozen=True) +class StoredProviderSecret: + """One encrypted upstream provider API key record.""" + + provider_name: str + secret_hint: str + created_at: int = 0 + updated_at: int = 0 + + def to_dict(self) -> dict[str, Any]: + return { + "provider_name": self.provider_name, + "configured": True, + "secret_hint": self.secret_hint, + "created_at": self.created_at, + "updated_at": self.updated_at, + "source": "dashboard", + } + + class GatewayAdminStore: """SQLite-backed dashboard users, sessions, and managed gateway tokens.""" @@ -159,6 +181,70 @@ def list_clients(self) -> list[dict[str, Any]]: ).fetchall() return [_stored_token_from_row(row).to_dict() for row in rows] + def provider_secret_storage_status(self) -> dict[str, Any]: + return _provider_secret_storage_status(self.config) + + def list_provider_secrets(self) -> dict[str, dict[str, Any]]: + with self._lock: + rows = self._conn().execute( + "SELECT * FROM gateway_provider_secrets ORDER BY provider_name" + ).fetchall() + return { + str(row["provider_name"]): _stored_provider_secret_from_row(row).to_dict() + for row in rows + } + + def provider_secret(self, provider_name: str) -> Optional[str]: + fernet = _provider_secret_fernet(self.config) + if fernet is None: + return None + with self._lock: + row = self._conn().execute( + "SELECT encrypted_api_key FROM gateway_provider_secrets WHERE provider_name = ?", + (provider_name,), + ).fetchone() + if row is None: + return None + try: + return fernet.decrypt(str(row["encrypted_api_key"]).encode("ascii")).decode("utf-8") + except Exception: + return None + + def set_provider_secret(self, provider_name: str, api_key: str) -> dict[str, Any]: + provider_name = str(provider_name or "").strip() + api_key = str(api_key or "").strip() + if not provider_name: + raise ValueError("Provider name is required") + if not api_key: + raise ValueError("Provider API key is required") + fernet = _require_provider_secret_fernet(self.config) + encrypted = fernet.encrypt(api_key.encode("utf-8")).decode("ascii") + now = int(time.time()) + with self._lock: + self._conn().execute( + """ + INSERT INTO gateway_provider_secrets ( + provider_name, encrypted_api_key, secret_hint, created_at, updated_at + ) + VALUES (?, ?, ?, ?, ?) + ON CONFLICT(provider_name) DO UPDATE SET + encrypted_api_key = excluded.encrypted_api_key, + secret_hint = excluded.secret_hint, + updated_at = excluded.updated_at + """, + (provider_name, encrypted, _secret_hint(api_key), now, now), + ) + self._conn().commit() + return self.list_provider_secrets()[provider_name] + + def delete_provider_secret(self, provider_name: str) -> None: + with self._lock: + self._conn().execute( + "DELETE FROM gateway_provider_secrets WHERE provider_name = ?", + (provider_name,), + ) + self._conn().commit() + def get_client(self, client_id: str) -> Optional[dict[str, Any]]: with self._lock: row = self._conn().execute( @@ -190,6 +276,7 @@ def client_from_token(self, token: str) -> Optional[GatewayClient]: team_id=stored.team_id, user_id=stored.user_id, allowed_models=stored.allowed_models, + policy_actions=normalize_policy_actions(stored.policy_actions), max_requests=stored.max_requests, max_tokens=stored.max_tokens, max_budget=stored.max_budget, @@ -204,6 +291,7 @@ def create_client( team_id: Optional[str] = None, user_id: Optional[str] = None, allowed_models: Optional[list[str]] = None, + policy_actions: Optional[dict[str, str]] = None, max_requests: Optional[int] = None, max_tokens: Optional[int] = None, max_budget: Optional[float] = None, @@ -216,17 +304,18 @@ def create_client( client_id = _client_id() now = int(time.time()) allowed = _normalize_allowed_models(allowed_models) + policy = normalize_policy_actions(policy_actions) with self._lock: try: self._conn().execute( """ INSERT INTO gateway_clients ( id, name, token_hash, token_prefix, enabled, - tenant_id, team_id, user_id, allowed_models_json, + tenant_id, team_id, user_id, allowed_models_json, policy_actions_json, max_requests, max_tokens, max_budget, budget_reset, created_at, updated_at, rotated_at, last_used_at ) - VALUES (?, ?, ?, ?, 1, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL) + VALUES (?, ?, ?, ?, 1, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL) """, ( client_id, @@ -237,6 +326,7 @@ def create_client( _blank_to_none(team_id), _blank_to_none(user_id), json.dumps(allowed, sort_keys=True), + json.dumps(policy, sort_keys=True), max_requests, max_tokens, max_budget, @@ -283,6 +373,7 @@ def update_client( team_id: Optional[str] = None, user_id: Optional[str] = None, allowed_models: Optional[list[str]] = None, + policy_actions: Optional[dict[str, str]] = None, max_requests: Optional[int] = None, max_tokens: Optional[int] = None, max_budget: Optional[float] = None, @@ -306,6 +397,11 @@ def update_client( _normalize_allowed_models(allowed_models), sort_keys=True, ) + if policy_actions is not None: + updates["policy_actions_json"] = json.dumps( + normalize_policy_actions(policy_actions), + sort_keys=True, + ) if max_requests is not None: updates["max_requests"] = max_requests if max_tokens is not None: @@ -427,6 +523,7 @@ def _init_db(self) -> None: team_id TEXT, user_id TEXT, allowed_models_json TEXT NOT NULL, + policy_actions_json TEXT NOT NULL DEFAULT '{}', max_requests INTEGER, max_tokens INTEGER, max_budget REAL, @@ -438,8 +535,34 @@ def _init_db(self) -> None: ) """ ) + self._conn().execute( + """ + CREATE TABLE IF NOT EXISTS gateway_provider_secrets ( + provider_name TEXT PRIMARY KEY, + encrypted_api_key TEXT NOT NULL, + secret_hint TEXT NOT NULL, + created_at INTEGER NOT NULL, + updated_at INTEGER NOT NULL + ) + """ + ) + self._ensure_column( + "gateway_clients", + "policy_actions_json", + "TEXT NOT NULL DEFAULT '{}'", + ) self._conn().commit() + def _ensure_column(self, table_name: str, column_name: str, definition: str) -> None: + columns = { + str(row["name"]) + for row in self._conn().execute(f"PRAGMA table_info({table_name})") + } + if column_name not in columns: + self._conn().execute( + f"ALTER TABLE {table_name} ADD COLUMN {column_name} {definition}" + ) + def _conn(self) -> sqlite3.Connection: conn = getattr(self, "_connection", None) if conn is None: @@ -470,6 +593,7 @@ def _admin_state_path(config: GatewayConfig) -> str: def _stored_token_from_row(row: sqlite3.Row) -> StoredGatewayToken: allowed_models = tuple(json.loads(row["allowed_models_json"] or "[]")) + policy_actions = normalize_policy_actions(json.loads(row["policy_actions_json"] or "{}")) return StoredGatewayToken( id=str(row["id"]), name=str(row["name"]), @@ -479,6 +603,7 @@ def _stored_token_from_row(row: sqlite3.Row) -> StoredGatewayToken: team_id=row["team_id"], user_id=row["user_id"], allowed_models=allowed_models, + policy_actions=policy_actions, max_requests=row["max_requests"], max_tokens=row["max_tokens"], max_budget=row["max_budget"], @@ -490,6 +615,15 @@ def _stored_token_from_row(row: sqlite3.Row) -> StoredGatewayToken: ) +def _stored_provider_secret_from_row(row: sqlite3.Row) -> StoredProviderSecret: + return StoredProviderSecret( + provider_name=str(row["provider_name"]), + secret_hint=str(row["secret_hint"]), + created_at=int(row["created_at"]), + updated_at=int(row["updated_at"]), + ) + + def _normalize_allowed_models(values: Optional[list[str]]) -> list[str]: normalized = [] for value in values or ["*"]: @@ -536,6 +670,63 @@ def _verify_password(password: str, stored_hash: str) -> bool: return False +def _provider_secret_storage_status(config: GatewayConfig) -> dict[str, Any]: + key_env = _provider_secret_key_env(config) + if not getattr(config, "provider_secret_storage_enabled", True): + return { + "enabled": False, + "key_env": key_env, + "reason": "dashboard provider secret storage is disabled", + } + try: + import cryptography.fernet # noqa: F401 + except ImportError: + return { + "enabled": False, + "key_env": key_env, + "reason": "install cryptography to store encrypted provider secrets", + } + if not os.getenv(key_env): + return { + "enabled": False, + "key_env": key_env, + "reason": f"set {key_env} to enable encrypted dashboard provider secrets", + } + return {"enabled": True, "key_env": key_env, "reason": "encrypted provider secret storage enabled"} + + +def _provider_secret_fernet(config: GatewayConfig) -> Any: + if not _provider_secret_storage_status(config)["enabled"]: + return None + from cryptography.fernet import Fernet + + raw_key = os.getenv(_provider_secret_key_env(config), "").strip().encode("utf-8") + try: + return Fernet(raw_key) + except Exception: + derived = base64.urlsafe_b64encode(hashlib.sha256(raw_key).digest()) + return Fernet(derived) + + +def _require_provider_secret_fernet(config: GatewayConfig) -> Any: + fernet = _provider_secret_fernet(config) + if fernet is None: + status = _provider_secret_storage_status(config) + raise ValueError(str(status["reason"])) + return fernet + + +def _provider_secret_key_env(config: GatewayConfig) -> str: + return getattr(config, "provider_secret_key_env", "SENTINELGUARD_ENCRYPTION_KEY") + + +def _secret_hint(secret: str) -> str: + text = str(secret or "").strip() + if len(text) < 4: + return "configured" + return f"configured ...{text[-4:]}" + + def _token_prefix(token: str) -> str: return f"{token[:12]}...{token[-4:]}" diff --git a/sentinelguard/gateway/config.py b/sentinelguard/gateway/config.py index e994289..b621964 100644 --- a/sentinelguard/gateway/config.py +++ b/sentinelguard/gateway/config.py @@ -9,6 +9,29 @@ import yaml +POLICY_ACTION_ALIASES = { + "audit_only": "audit", + "log": "audit", + "log_only": "audit", + "mask": "redact", + "monitor": "audit", + "sanitize": "redact", + "warn": "audit", +} +POLICY_CATEGORY_ALIASES = { + "jailbreak": "attack", + "jailbreaks": "attack", + "prompt_attack": "attack", + "prompt_attacks": "attack", + "prompt_injection": "attack", + "prompt_injections": "attack", + "secret": "secret", + "secrets": "secret", +} +ALLOWED_POLICY_ACTIONS = {"allow", "audit", "block", "redact"} +ALLOWED_POLICY_CATEGORIES = {"attack", "secret", "pii", "pci", "phi", "other"} + + @dataclass class ProviderConfig: """One upstream model provider candidate for gateway routing.""" @@ -44,7 +67,7 @@ def to_dict(self) -> Dict[str, Any]: "upstream_model": self.upstream_model, "upstream_url": self.upstream_url, "api_key_env": self.api_key_env, - "api_key": self.api_key, + "api_key": "" if self.api_key else None, "enabled": self.enabled, "private": self.private, "priority": self.priority, @@ -70,6 +93,7 @@ class VirtualKeyConfig: team_id: Optional[str] = None user_id: Optional[str] = None allowed_models: List[str] = field(default_factory=list) + policy_actions: Dict[str, str] = field(default_factory=dict) max_requests: Optional[int] = None max_tokens: Optional[int] = None max_budget: Optional[float] = None @@ -78,7 +102,11 @@ class VirtualKeyConfig: @classmethod def from_dict(cls, data: Dict[str, Any]) -> VirtualKeyConfig: known_fields = cls.__dataclass_fields__ - return cls(**{key: value for key, value in data.items() if key in known_fields}) + values = {key: value for key, value in data.items() if key in known_fields} + policy_actions = data.get("policy_actions", data.get("policy")) + if policy_actions is not None: + values["policy_actions"] = normalize_policy_actions(policy_actions) + return cls(**values) def to_dict(self) -> Dict[str, Any]: return { @@ -90,6 +118,7 @@ def to_dict(self) -> Dict[str, Any]: "team_id": self.team_id, "user_id": self.user_id, "allowed_models": list(self.allowed_models), + "policy_actions": normalize_policy_actions(self.policy_actions), "max_requests": self.max_requests, "max_tokens": self.max_tokens, "max_budget": self.max_budget, @@ -188,6 +217,58 @@ def _list_value(value: Any) -> List[str]: return [str(value).strip()] if str(value).strip() else [] +def normalize_policy_actions(value: Any) -> Dict[str, str]: + """Normalize per-client gateway policy actions. + + Accepted categories are attack, secret, pii, pci, phi, and other. Actions + are allow, audit, block, and redact. Common aliases such as secrets, + prompt_attack, sanitize, mask, warn, and log_only are accepted. + """ + if value is None: + return {} + if isinstance(value, str): + return _policy_actions_from_string(value) + if not isinstance(value, dict): + return {} + + normalized: Dict[str, str] = {} + for raw_key, raw_action in value.items(): + category = _normalize_policy_category(raw_key) + action = _normalize_policy_action(raw_action) + if category and action: + normalized[category] = action + return normalized + + +def _policy_actions_from_string(value: str) -> Dict[str, str]: + items = [item.strip() for item in value.split(",") if item.strip()] + parsed: Dict[str, str] = {} + for item in items: + if ":" in item: + key, action = item.split(":", 1) + elif "=" in item: + key, action = item.split("=", 1) + else: + continue + category = _normalize_policy_category(key) + normalized_action = _normalize_policy_action(action) + if category and normalized_action: + parsed[category] = normalized_action + return parsed + + +def _normalize_policy_category(value: Any) -> Optional[str]: + normalized = str(value or "").strip().lower().replace("-", "_") + normalized = POLICY_CATEGORY_ALIASES.get(normalized, normalized) + return normalized if normalized in ALLOWED_POLICY_CATEGORIES else None + + +def _normalize_policy_action(value: Any) -> Optional[str]: + normalized = str(value or "").strip().lower().replace("-", "_") + normalized = POLICY_ACTION_ALIASES.get(normalized, normalized) + return normalized if normalized in ALLOWED_POLICY_ACTIONS else None + + def _normalize_guardrail_mode(value: Any) -> str: normalized = str(value or "enforce").strip().lower().replace("-", "_") if normalized in {"log", "log_only", "logging"}: @@ -264,6 +345,8 @@ class GatewayConfig: admin_password_env: str = "SENTINELGUARD_ADMIN_PASSWORD" admin_viewer_username_env: str = "SENTINELGUARD_VIEWER_USERNAME" admin_viewer_password_env: str = "SENTINELGUARD_VIEWER_PASSWORD" + provider_secret_storage_enabled: bool = True + provider_secret_key_env: str = "SENTINELGUARD_ENCRYPTION_KEY" otel_enabled: bool = False langfuse_enabled: bool = False metrics_enabled: bool = True @@ -325,7 +408,7 @@ def to_dict(self) -> Dict[str, Any]: "provider": self.provider, "upstream_url": self.upstream_url, "api_key_env": self.api_key_env, - "api_key": self.api_key, + "api_key": "" if self.api_key else None, "providers": [provider.to_dict() for provider in self.providers], "virtual_keys": [virtual_key.to_dict() for virtual_key in self.virtual_keys], "client_api_key_env": self.client_api_key_env, @@ -376,6 +459,8 @@ def to_dict(self) -> Dict[str, Any]: "admin_password_env": self.admin_password_env, "admin_viewer_username_env": self.admin_viewer_username_env, "admin_viewer_password_env": self.admin_viewer_password_env, + "provider_secret_storage_enabled": self.provider_secret_storage_enabled, + "provider_secret_key_env": self.provider_secret_key_env, "otel_enabled": self.otel_enabled, "langfuse_enabled": self.langfuse_enabled, "metrics_enabled": self.metrics_enabled, diff --git a/sentinelguard/gateway/guardrails.py b/sentinelguard/gateway/guardrails.py index 4fa8bdf..4b537c9 100644 --- a/sentinelguard/gateway/guardrails.py +++ b/sentinelguard/gateway/guardrails.py @@ -162,6 +162,7 @@ async def apply_gateway_guardrails( prompt: Optional[str] = None, requested_guardrails: Sequence[str] = (), metric_direction: Optional[str] = None, + client: Optional[GatewayClient] = None, streaming: bool = False, ) -> GuardrailApplyResult: """Scan text once and apply all active named guardrails for the stage.""" @@ -196,7 +197,13 @@ async def apply_gateway_guardrails( started_at = time.perf_counter() scan = await _scan_text(guard, text, normalized_direction, prompt=prompt) record_scan(metric_direction or normalized_direction, scan) - base_decision = _evaluate_text_policy(text, scan, config, normalized_direction) + base_decision = _evaluate_text_policy( + text, + scan, + config, + normalized_direction, + client=client, + ) elapsed_ms = (time.perf_counter() - started_at) * 1000 effective_decision = _effective_decision(base_decision, active, normalized_direction) executions = tuple( @@ -416,10 +423,13 @@ def _evaluate_text_policy( scan: AggregatedResult, config: GatewayConfig, direction: str, + *, + client: Optional[GatewayClient] = None, ) -> PolicyDecision: + client_policy = client.policy_actions if client is not None else None if direction == "output": - return evaluate_output_policy(text, scan, config) - return evaluate_prompt_policy(text, scan, config) + return evaluate_output_policy(text, scan, config, client_policy=client_policy) + return evaluate_prompt_policy(text, scan, config, client_policy=client_policy) def _scan_summary(scan: AggregatedResult) -> dict[str, Any]: diff --git a/sentinelguard/gateway/operations.py b/sentinelguard/gateway/operations.py index b83471e..f4d1101 100644 --- a/sentinelguard/gateway/operations.py +++ b/sentinelguard/gateway/operations.py @@ -19,7 +19,12 @@ from pathlib import Path from typing import Any, Mapping, Optional -from sentinelguard.gateway.config import GatewayConfig, ProviderConfig, VirtualKeyConfig +from sentinelguard.gateway.config import ( + GatewayConfig, + ProviderConfig, + VirtualKeyConfig, + normalize_policy_actions, +) from sentinelguard.gateway.providers import extract_assistant_text, extract_last_user_text @@ -33,6 +38,7 @@ class GatewayClient: team_id: Optional[str] = None user_id: Optional[str] = None allowed_models: tuple[str, ...] = () + policy_actions: Mapping[str, str] = field(default_factory=dict) max_requests: Optional[int] = None max_tokens: Optional[int] = None max_budget: Optional[float] = None @@ -715,6 +721,7 @@ def virtual_key_summary(config: GatewayConfig) -> list[dict[str, Any]]: "team_id": key.team_id, "user_id": key.user_id, "allowed_models": list(key.allowed_models), + "policy_actions": normalize_policy_actions(key.policy_actions), "max_requests": key.max_requests, "max_tokens": key.max_tokens, "max_budget": key.max_budget, @@ -754,6 +761,7 @@ def _client_from_virtual_key(virtual_key: VirtualKeyConfig, secret: str) -> Gate team_id=virtual_key.team_id, user_id=virtual_key.user_id, allowed_models=tuple(virtual_key.allowed_models), + policy_actions=normalize_policy_actions(virtual_key.policy_actions), max_requests=virtual_key.max_requests, max_tokens=virtual_key.max_tokens, max_budget=virtual_key.max_budget, diff --git a/sentinelguard/gateway/policy.py b/sentinelguard/gateway/policy.py index 0971922..98c0a7f 100644 --- a/sentinelguard/gateway/policy.py +++ b/sentinelguard/gateway/policy.py @@ -4,7 +4,7 @@ from dataclasses import dataclass, field from enum import Enum -from typing import List, Optional +from typing import List, Mapping, Optional from sentinelguard.core.scanner import AggregatedResult, ScanResult from sentinelguard.gateway.config import GatewayConfig @@ -15,6 +15,7 @@ class PolicyAction(str, Enum): """Actions the gateway can take after scanner evaluation.""" ALLOW = "allow" + AUDIT = "audit" REDACT = "redact" BLOCK = "block" REVIEW = "review" @@ -35,13 +36,14 @@ class PolicyDecision: @property def allowed(self) -> bool: - return self.action in {PolicyAction.ALLOW, PolicyAction.REDACT} + return self.action in {PolicyAction.ALLOW, PolicyAction.AUDIT, PolicyAction.REDACT} def evaluate_prompt_policy( text: str, result: AggregatedResult, config: GatewayConfig, + client_policy: Optional[Mapping[str, str]] = None, ) -> PolicyDecision: """Convert prompt scan results into an explicit gateway decision.""" return _evaluate_policy( @@ -51,6 +53,7 @@ def evaluate_prompt_policy( direction="prompt", block_on_fail=config.block_on_prompt_fail, redact_pii=config.redact_pii, + client_policy=client_policy, ) @@ -58,6 +61,7 @@ def evaluate_output_policy( text: str, result: AggregatedResult, config: GatewayConfig, + client_policy: Optional[Mapping[str, str]] = None, ) -> PolicyDecision: """Convert output scan results into an explicit gateway decision.""" return _evaluate_policy( @@ -67,6 +71,7 @@ def evaluate_output_policy( direction="output", block_on_fail=config.block_on_output_fail, redact_pii=config.redact_output_pii, + client_policy=client_policy, ) @@ -78,6 +83,7 @@ def _evaluate_policy( direction: str, block_on_fail: bool, redact_pii: bool, + client_policy: Optional[Mapping[str, str]], ) -> PolicyDecision: categories = _failed_categories(result) reason_codes = _reason_codes(result) @@ -95,17 +101,88 @@ def _evaluate_policy( ) sanitized_text = result.sanitized_output - if redact_pii and "pii" in categories: + actions = _category_actions( + categories, + result, + client_policy, + block_on_fail=block_on_fail, + redact_pii=redact_pii, + ) + if any(action == PolicyAction.REDACT for action in actions.values()): if sanitized_text and sanitized_text != text: - reason_codes.append("pii_redacted") + reason_codes.append(_redaction_reason(categories)) else: - redacted_text = _redact_pii(text) + redacted_text = _redact_sensitive_text(text, categories) if redacted_text != text: sanitized_text = redacted_text - reason_codes.append("pii_redacted") + reason_codes.append(_redaction_reason(categories)) else: sanitized_text = None - reason_codes.append("pii_redaction_unavailable") + reason_codes.append("redaction_unavailable") + + reason_codes.extend( + f"policy:{category}_{action.value}" + for category, action in sorted(actions.items()) + if action != PolicyAction.ALLOW + ) + + if any(action == PolicyAction.BLOCK for action in actions.values()): + return PolicyDecision( + action=PolicyAction.BLOCK, + direction=direction, + reason_codes=_dedupe(reason_codes), + failed_scanners=list(result.failed_scanners), + warning_scanners=list(result.warning_scanners), + route_constraint=route_constraint, + highest_risk=result.highest_risk.value, + ) + + if sanitized_text and sanitized_text != text and any( + action == PolicyAction.REDACT for action in actions.values() + ): + return PolicyDecision( + action=PolicyAction.REDACT, + direction=direction, + reason_codes=_dedupe(reason_codes), + failed_scanners=list(result.failed_scanners), + warning_scanners=list(result.warning_scanners), + sanitized_text=sanitized_text, + route_constraint=route_constraint, + highest_risk=result.highest_risk.value, + ) + + if any(action == PolicyAction.REDACT for action in actions.values()): + return PolicyDecision( + action=PolicyAction.BLOCK, + direction=direction, + reason_codes=_dedupe(reason_codes or ["redaction_unavailable"]), + failed_scanners=list(result.failed_scanners), + warning_scanners=list(result.warning_scanners), + route_constraint=route_constraint, + highest_risk=result.highest_risk.value, + ) + + if any(action == PolicyAction.AUDIT for action in actions.values()): + return PolicyDecision( + action=PolicyAction.AUDIT, + direction=direction, + reason_codes=_dedupe(reason_codes), + failed_scanners=list(result.failed_scanners), + warning_scanners=list(result.warning_scanners), + route_constraint=route_constraint, + highest_risk=result.highest_risk.value, + ) + + if categories and all(action == PolicyAction.ALLOW for action in actions.values()): + return PolicyDecision( + action=PolicyAction.ALLOW, + direction=direction, + reason_codes=_dedupe(reason_codes), + failed_scanners=list(result.failed_scanners), + warning_scanners=list(result.warning_scanners), + route_constraint=route_constraint, + highest_risk=result.highest_risk.value, + ) if not result.is_valid and block_on_fail: if "pii" in categories and categories <= {"pii"} and sanitized_text: @@ -156,10 +233,88 @@ def _failed_categories(result: AggregatedResult) -> set[str]: categories = set() for scan in result.results: if not scan.is_valid: - categories.add(scanner_category(scan.scanner_name)) + category = scanner_category(scan.scanner_name) + categories.add(category) + if category == "pii": + categories.update(_sensitive_data_subcategories(scan)) return categories +def _category_actions( + categories: set[str], + result: AggregatedResult, + client_policy: Optional[Mapping[str, str]], + *, + block_on_fail: bool, + redact_pii: bool, +) -> dict[str, PolicyAction]: + policy = _normalize_client_policy(client_policy) + return { + category: _policy_action_for_category( + category, + result, + policy, + block_on_fail=block_on_fail, + redact_pii=redact_pii, + ) + for category in categories + } + + +def _policy_action_for_category( + category: str, + result: AggregatedResult, + policy: Mapping[str, PolicyAction], + *, + block_on_fail: bool, + redact_pii: bool, +) -> PolicyAction: + if category in policy: + return policy[category] + if category in {"pii", "pci", "phi"} and redact_pii: + return PolicyAction.REDACT + if result.is_valid: + return PolicyAction.ALLOW + return PolicyAction.BLOCK if block_on_fail else PolicyAction.ALLOW + + +def _normalize_client_policy( + policy: Optional[Mapping[str, str]], +) -> dict[str, PolicyAction]: + normalized = {} + for raw_category, raw_action in (policy or {}).items(): + category = str(raw_category or "").strip().lower().replace("-", "_") + category = {"secrets": "secret", "prompt_attack": "attack"}.get(category, category) + action = str(raw_action or "").strip().lower().replace("-", "_") + action = { + "audit_only": "audit", + "log": "audit", + "log_only": "audit", + "mask": "redact", + "monitor": "audit", + "sanitize": "redact", + "warn": "audit", + }.get(action, action) + if category in {"attack", "secret", "pii", "pci", "phi", "other"} and action in { + "allow", + "audit", + "block", + "redact", + }: + normalized[category] = PolicyAction(action) + return normalized + + +def _sensitive_data_subcategories(scan: ScanResult) -> set[str]: + entity_types = {str(entity).upper() for entity in scan.details.get("entity_types", []) or []} + found = set() + if entity_types & {"CREDIT_CARD", "IBAN_CODE", "US_BANK_NUMBER", "CRYPTO"}: + found.add("pci") + if entity_types & {"MEDICAL_LICENSE", "UK_NHS", "US_NPI"}: + found.add("phi") + return found + + def _reason_codes(result: AggregatedResult) -> List[str]: codes = [] for scan in result.results: @@ -186,5 +341,17 @@ def _redact_pii(text: str) -> str: return text +def _redact_sensitive_text(text: str, categories: set[str]) -> str: + if categories & {"pii", "pci", "phi"}: + return _redact_pii(text) + return text + + +def _redaction_reason(categories: set[str]) -> str: + if "pii" in categories: + return "pii_redacted" + return "content_redacted" + + def _dedupe(values: List[str]) -> List[str]: return list(dict.fromkeys(values)) diff --git a/sentinelguard/gateway/providers.py b/sentinelguard/gateway/providers.py index 04d9899..236bfe2 100644 --- a/sentinelguard/gateway/providers.py +++ b/sentinelguard/gateway/providers.py @@ -382,8 +382,9 @@ def effective_provider(config: GatewayConfig) -> str: def configured_providers(config: GatewayConfig) -> list[ProviderConfig]: """Return configured provider candidates, preserving single-provider compatibility.""" if config.providers: - return [provider for provider in config.providers if provider.enabled] - return [ + providers = [provider for provider in config.providers if provider.enabled] + return _with_dashboard_provider_secrets(config, providers) + providers = [ ProviderConfig( name=effective_provider(config), provider=config.provider, @@ -399,6 +400,26 @@ def configured_providers(config: GatewayConfig) -> list[ProviderConfig]: timeout_seconds=config.timeout_seconds, ) ] + return _with_dashboard_provider_secrets(config, providers) + + +def _with_dashboard_provider_secrets( + config: GatewayConfig, + providers: list[ProviderConfig], +) -> list[ProviderConfig]: + if not config.admin_ui_enabled or not getattr(config, "provider_secret_storage_enabled", True): + return providers + try: + from sentinelguard.gateway.admin import gateway_admin_store + + admin_store = gateway_admin_store(config) + updated = [] + for provider in providers: + dashboard_key = admin_store.provider_secret(provider.name) + updated.append(replace(provider, api_key=dashboard_key) if dashboard_key else provider) + return updated + except Exception: + return providers def select_provider_sequence( @@ -720,6 +741,15 @@ def effective_api_key_env(config: GatewayConfig) -> str: return configured or "OPENAI_API_KEY" +def api_key_env_names_for_provider( + config: GatewayConfig, + provider: Optional[ProviderConfig] = None, +) -> list[str]: + """Return all environment variable names checked for a provider key.""" + provider_config = _gateway_config_for_provider(config, provider) if provider else config + return _api_key_env_names(provider_config) + + async def _forward_openai_compatible( httpx: Any, payload: Mapping[str, Any], diff --git a/sentinelguard/gateway/server.py b/sentinelguard/gateway/server.py index c63c647..62fa5a4 100644 --- a/sentinelguard/gateway/server.py +++ b/sentinelguard/gateway/server.py @@ -20,7 +20,7 @@ GatewayAdminStore, gateway_admin_store, ) -from sentinelguard.gateway.config import GatewayConfig +from sentinelguard.gateway.config import GatewayConfig, normalize_policy_actions from sentinelguard.gateway.guardrails import ( apply_gateway_guardrails, guardrail_summary, @@ -44,6 +44,7 @@ ) from sentinelguard.gateway.observability import emit_gateway_event from sentinelguard.gateway.providers import ( + api_key_env_names_for_provider, available_gateway_models, configured_providers, effective_api_key_env, @@ -201,6 +202,7 @@ async def guardrails_apply(request: Request): stage=str(payload.get("stage") or "manual"), prompt=_payload_text(payload, "prompt"), requested_guardrails=requested, + client=auth.client, ) return JSONResponse(content=result.to_dict(include_text=True), status_code=200) @@ -271,6 +273,7 @@ async def admin_create_client(request: Request): team_id=_payload_text(payload, "team_id"), user_id=_payload_text(payload, "user_id"), allowed_models=_payload_allowed_models(payload.get("allowed_models")), + policy_actions=_payload_policy_actions(payload.get("policy_actions")), max_requests=_optional_int(payload.get("max_requests")), max_tokens=_optional_int(payload.get("max_tokens")), max_budget=_optional_float(payload.get("max_budget")), @@ -310,6 +313,11 @@ async def admin_update_client(client_id: str, request: Request): if "allowed_models" in payload else None ), + policy_actions=( + _payload_policy_actions(payload.get("policy_actions")) + if "policy_actions" in payload + else None + ), max_requests=_optional_int(payload.get("max_requests")), max_tokens=_optional_int(payload.get("max_tokens")), max_budget=_optional_float(payload.get("max_budget")), @@ -352,6 +360,45 @@ async def admin_client_usage(client_id: str, request: Request): _require_admin_user(request, admin_store, config) return _client_usage_payload(client_id, config, admin_store, usage_store) + @app.get("/admin/api/providers") + async def admin_providers(request: Request): + _require_admin_user(request, admin_store, config) + return _admin_provider_secrets_payload(config, admin_store) + + @app.patch("/admin/api/providers/{provider_name}/secret") + async def admin_update_provider_secret(provider_name: str, request: Request): + _require_admin_user(request, admin_store, config, required_role=ADMIN_ROLE) + if provider_name not in _configured_provider_names(config): + return _gateway_error_response( + 404, + "sentinelguard_provider_not_found", + "SentinelGuard provider route not found", + ) + payload = await _json_body(request) + try: + provider_secret = admin_store.set_provider_secret( + provider_name, + str(payload.get("api_key") or payload.get("secret") or ""), + ) + except ValueError as exc: + return _gateway_error_response(400, "sentinelguard_provider_secret_error", str(exc)) + return { + "provider_secret": provider_secret, + "providers": _admin_provider_secrets_payload(config, admin_store)["providers"], + } + + @app.delete("/admin/api/providers/{provider_name}/secret") + async def admin_delete_provider_secret(provider_name: str, request: Request): + _require_admin_user(request, admin_store, config, required_role=ADMIN_ROLE) + if provider_name not in _configured_provider_names(config): + return _gateway_error_response( + 404, + "sentinelguard_provider_not_found", + "SentinelGuard provider route not found", + ) + admin_store.delete_provider_secret(provider_name) + return _admin_provider_secrets_payload(config, admin_store) + @app.post("/gateway/v1/client/token/rotate") async def rotate_current_client_token(request: Request): auth = authenticate_gateway_request(request.headers, config) @@ -495,6 +542,7 @@ async def chat_completions(request: Request): direction="prompt", stage="pre_call", requested_guardrails=requested_guardrails, + client=auth.client, streaming=False, ) prompt_scan = prompt_guardrail.scan @@ -549,6 +597,7 @@ async def chat_completions(request: Request): stage="post_call", prompt=safe_prompt, requested_guardrails=requested_guardrails, + client=auth.client, streaming=False, ) output_scan = output_guardrail.scan @@ -594,6 +643,7 @@ async def chat_completions(request: Request): stage="post_call", prompt=safe_prompt, requested_guardrails=requested_guardrails, + client=auth.client, streaming=False, ) output_scan = output_guardrail.scan @@ -687,6 +737,7 @@ async def _handle_streaming_chat( direction="prompt", stage="pre_call", requested_guardrails=requested_guardrails, + client=client, streaming=True, ) prompt_scan = prompt_guardrail.scan @@ -749,6 +800,7 @@ async def _handle_streaming_chat( stage="post_call", prompt=safe_prompt, requested_guardrails=requested_guardrails, + client=client, streaming=True, ) output_scan = output_guardrail.scan @@ -815,6 +867,7 @@ async def _passthrough_gateway_request( stage="passthrough", requested_guardrails=requested, metric_direction=gateway_name, + client=auth.client, ) scan = guardrail_result.scan decision = guardrail_result.decision @@ -1087,6 +1140,7 @@ def _usage_payload(client: GatewayClient, usage_store: Any) -> dict[str, Any]: "tenant_id": client.tenant_id, "team_id": client.team_id, "user_id": client.user_id, + "policy_actions": normalize_policy_actions(client.policy_actions), }, "usage": usage_store.snapshot( client.key_id, @@ -1369,6 +1423,12 @@ def _admin_html() -> str: + + + + + +

@@ -1393,6 +1453,12 @@ def _admin_html() -> str: + + + + + + @@ -1402,7 +1468,7 @@ def _admin_html() -> str:

All Clients

- +
NameSourceStatusModelsRequestsLast used
NameSourceStatusModelsPolicyRequestsLast used
@@ -1410,19 +1476,40 @@ def _admin_html() -> str:

Provider Health

+
- +
ProviderModelStatusAttemptsSuccessesFailures
ProviderModelKey sourceSecretStatusAttemptsSuccessesFailures
+
+

Update Provider Key

+
+ + +
+
+ + +
+

+
@@ -1768,6 +1946,7 @@ def _admin_summary_payload( "summary": summary, "clients": clients, "provider_health": _provider_health_payload(config)["providers"], + "provider_secrets": _admin_provider_secrets_payload(config, admin_store), "routes": _routes_payload(config), "security_warnings": admin_store.security_warnings(), "metrics": { @@ -1804,6 +1983,7 @@ def _configured_client_summaries(config: GatewayConfig, usage_store: Any) -> lis "team_id": key.team_id, "user_id": key.user_id, "allowed_models": list(key.allowed_models), + "policy_actions": normalize_policy_actions(key.policy_actions), "max_requests": key.max_requests, "max_tokens": key.max_tokens, "max_budget": key.max_budget, @@ -1833,6 +2013,7 @@ def _configured_client_summaries(config: GatewayConfig, usage_store: Any) -> lis "team_id": None, "user_id": None, "allowed_models": ["*"], + "policy_actions": {}, "max_requests": None, "max_tokens": None, "max_budget": None, @@ -1850,6 +2031,86 @@ def _configured_client_summaries(config: GatewayConfig, usage_store: Any) -> lis return clients +def _admin_provider_secrets_payload( + config: GatewayConfig, + admin_store: GatewayAdminStore, +) -> dict[str, Any]: + secret_map = admin_store.list_provider_secrets() + storage_status = admin_store.provider_secret_storage_status() + providers = [] + for provider in configured_providers(config): + dashboard_secret = secret_map.get(provider.name) + dashboard_secret_usable = bool( + dashboard_secret + and storage_status.get("enabled") + and admin_store.provider_secret(provider.name) + ) + providers.append( + _provider_secret_summary(config, provider, dashboard_secret, dashboard_secret_usable) + ) + return { + "object": "sentinelguard.gateway.admin.provider_secrets", + "api_version": GATEWAY_API_VERSION, + "secret_storage": storage_status, + "providers": providers, + } + + +def _provider_secret_summary( + config: GatewayConfig, + provider: Any, + dashboard_secret: Optional[Mapping[str, Any]], + dashboard_secret_usable: bool, +) -> dict[str, Any]: + env_names = api_key_env_names_for_provider(config, provider) + configured_env = next((env_name for env_name in env_names if os.getenv(env_name)), None) + env_configured = configured_env is not None + yaml_configured = bool(provider.api_key) + dashboard_secret_error = None + if dashboard_secret and dashboard_secret_usable: + api_key_source = "dashboard" + secret_hint = dashboard_secret.get("secret_hint") + configured = True + elif yaml_configured: + api_key_source = "yaml" + secret_hint = "configured in YAML" + configured = True + elif env_configured: + api_key_source = "env" + secret_hint = f"configured from {configured_env}" + configured = True + elif dashboard_secret: + api_key_source = "dashboard_unavailable" + secret_hint = "dashboard key cannot be decrypted without the encryption key" + dashboard_secret_error = "Dashboard key is stored but unavailable to runtime" + configured = False + else: + api_key_source = "missing" + secret_hint = "not configured" + configured = False + return { + "name": provider.name, + "provider": provider.provider, + "model_name": provider.model_name, + "upstream_model": provider.upstream_model, + "upstream_url": provider.upstream_url, + "api_key_env": provider.api_key_env, + "api_key_env_names": env_names, + "api_key_configured": configured, + "api_key_source": api_key_source, + "secret_hint": secret_hint, + "dashboard_secret": bool(dashboard_secret), + "dashboard_secret_usable": dashboard_secret_usable, + "dashboard_secret_error": dashboard_secret_error, + "enabled": provider.enabled, + "private": provider.private, + } + + +def _configured_provider_names(config: GatewayConfig) -> set[str]: + return {provider.name for provider in configured_providers(config)} + + def _with_usage(client: Mapping[str, Any], usage_store: Any) -> dict[str, Any]: item = dict(client) item["usage"] = usage_store.snapshot( @@ -1962,6 +2223,10 @@ def _payload_allowed_models(value: Any) -> Optional[list[str]]: return None +def _payload_policy_actions(value: Any) -> dict[str, str]: + return normalize_policy_actions(value) + + def _optional_int(value: Any) -> Optional[int]: if value in (None, ""): return None diff --git a/sentinelguard/monitoring.py b/sentinelguard/monitoring.py index ed750ab..181ad95 100644 --- a/sentinelguard/monitoring.py +++ b/sentinelguard/monitoring.py @@ -211,7 +211,7 @@ def _safe_action(action: Optional[str]) -> str: if not action: return "block" normalized = action.lower() - if normalized in {"block", "warn", "allow", "sanitize", "redact", "review"}: + if normalized in {"audit", "block", "warn", "allow", "sanitize", "redact", "review"}: return normalized return "other" diff --git a/tests/test_cli.py b/tests/test_cli.py index 4dee147..a1a0eb6 100644 --- a/tests/test_cli.py +++ b/tests/test_cli.py @@ -77,6 +77,7 @@ def test_init_with_env_generates_local_tokens(tmp_path, capsys): env_text = env_file.read_text(encoding="utf-8") assert "SENTINELGUARD_GATEWAY_API_KEY=sgw_" in env_text assert "SENTINELGUARD_AUDIT_SALT=sgaudit_" in env_text + assert "SENTINELGUARD_ENCRYPTION_KEY=sgencrypt_" in env_text assert "OPENAI_API_KEY=" in env_text output = capsys.readouterr().out diff --git a/tests/test_gateway.py b/tests/test_gateway.py index 81af555..be9d10a 100644 --- a/tests/test_gateway.py +++ b/tests/test_gateway.py @@ -1,12 +1,14 @@ """Tests for SentinelGuard gateway helpers.""" import json +import sqlite3 import pytest from sentinelguard.core.config import GuardConfig from sentinelguard.core.guard import SentinelGuard from sentinelguard.core.scanner import AggregatedResult, RiskLevel, ScanResult +from sentinelguard.gateway.admin import gateway_admin_store from sentinelguard.gateway.config import ( ComplexityRouterConfig, GatewayConfig, @@ -29,6 +31,7 @@ from sentinelguard.gateway.policy import PolicyAction, evaluate_prompt_policy from sentinelguard.gateway.providers import ( available_gateway_models, + configured_providers, effective_api_key_env, effective_provider, effective_upstream_url, @@ -173,6 +176,11 @@ def test_from_dict_with_model_routes_and_virtual_keys(self): "key": "sg-test-key", "team_id": "team-a", "allowed_models": ["fast-chat"], + "policy_actions": { + "prompt_attack": "block", + "secrets": "block", + "pii": "redact", + }, "max_requests": 10, "budget_reset": "daily", } @@ -196,9 +204,19 @@ def test_from_dict_with_model_routes_and_virtual_keys(self): assert config.providers[0].upstream_model == "gpt-4o-mini" assert config.providers[0].max_parallel_requests == 5 assert config.virtual_keys[0].allowed_models == ["fast-chat"] + assert config.virtual_keys[0].policy_actions == { + "attack": "block", + "secret": "block", + "pii": "redact", + } assert config.virtual_keys[0].budget_reset == "daily" assert config.to_dict()["complexity_router"]["enabled"] is True assert config.to_dict()["virtual_keys"][0]["key"] == "" + assert config.to_dict()["virtual_keys"][0]["policy_actions"] == { + "attack": "block", + "secret": "block", + "pii": "redact", + } def test_from_dict_with_named_guardrails_and_sensitive_routing(self): config = GatewayConfig.from_dict( @@ -318,6 +336,7 @@ def test_virtual_key_auth_returns_client_metadata(self): "tenant_id": "tenant-1", "team_id": "research", "allowed_models": ["fast-chat"], + "policy_actions": {"pii": "audit", "secrets": "block"}, } ] } @@ -329,6 +348,7 @@ def test_virtual_key_auth_returns_client_metadata(self): assert auth.client is not None assert auth.client.name == "research-team" assert auth.client.team_id == "research" + assert auth.client.policy_actions == {"pii": "audit", "secret": "block"} assert not authenticate_gateway_request({"authorization": "Bearer bad"}, config).allowed @@ -559,6 +579,55 @@ def test_guardrails_apply_endpoint_rejects_unknown_guardrail(self): assert response.status_code == 400 assert response.json()["error"]["type"] == "sentinelguard_unknown_guardrail" + def test_guardrails_apply_uses_authenticated_client_policy(self, monkeypatch): + fastapi_testclient = pytest.importorskip("fastapi.testclient") + + async def fake_scan_prompt(self, text, **kwargs): + return AggregatedResult( + is_valid=False, + results=[ + ScanResult( + is_valid=False, + score=0.9, + risk_level=RiskLevel.HIGH, + scanner_name="pii", + sanitized_output="Contact ", + ) + ], + failed_scanners=["pii"], + scanner_actions={"pii": "block"}, + sanitized_output="Contact ", + ) + + monkeypatch.setattr(SentinelGuard, "scan_prompt_async", fake_scan_prompt) + app = create_gateway_app( + guard_config=GuardConfig.preset_empty(), + gateway_config=GatewayConfig.from_dict( + { + "virtual_keys": [ + { + "name": "pii-audit-client", + "key": "sg-pii-audit", + "policy_actions": {"pii": "audit"}, + } + ] + } + ), + ) + client = fastapi_testclient.TestClient(app) + + response = client.post( + "/gateway/v1/guardrails/apply", + headers={"authorization": "Bearer sg-pii-audit"}, + json={"input": "Contact jane@example.com"}, + ) + + assert response.status_code == 200 + body = response.json() + assert body["allowed"] is True + assert body["action"] == "audit" + assert "policy:pii_audit" in body["reason_codes"] + def test_logging_only_guardrail_does_not_block(self, monkeypatch): fastapi_testclient = pytest.importorskip("fastapi.testclient") @@ -737,6 +806,11 @@ def test_admin_can_create_client_token_and_client_can_rotate_it(self, tmp_path, "tenant_id": "tenant-a", "team_id": "platform", "allowed_models": "fast-chat", + "policy_actions": { + "attack": "block", + "secret": "block", + "pii": "redact", + }, }, ) assert created.status_code == 201 @@ -748,6 +822,7 @@ def test_admin_can_create_client_token_and_client_can_rotate_it(self, tmp_path, assert "editForm" in client.get("/admin").text assert "editModels" in client.get("/admin").text + assert "editPolicyPii" in client.get("/admin").text auth = authenticate_gateway_request({"authorization": f"Bearer {raw_token}"}, config) assert auth.allowed @@ -755,6 +830,11 @@ def test_admin_can_create_client_token_and_client_can_rotate_it(self, tmp_path, assert auth.client.name == "chatbot-prod" assert auth.client.team_id == "platform" assert auth.client.allowed_models == ("fast-chat",) + assert auth.client.policy_actions == { + "attack": "block", + "secret": "block", + "pii": "redact", + } updated = client.patch( f"/admin/api/clients/{client_id}", @@ -763,6 +843,13 @@ def test_admin_can_create_client_token_and_client_can_rotate_it(self, tmp_path, "tenant_id": "tenant-a", "team_id": "ai-platform", "user_id": "chatbot-service", + "policy_actions": { + "attack": "block", + "secret": "block", + "pii": "audit", + "pci": "block", + "phi": "redact", + }, }, ) assert updated.status_code == 200 @@ -774,6 +861,13 @@ def test_admin_can_create_client_token_and_client_can_rotate_it(self, tmp_path, "private-chat", ] assert updated_client["team_id"] == "ai-platform" + assert updated_client["policy_actions"] == { + "attack": "block", + "secret": "block", + "pii": "audit", + "pci": "block", + "phi": "redact", + } updated_auth = authenticate_gateway_request({"authorization": f"Bearer {raw_token}"}, config) assert updated_auth.allowed @@ -786,6 +880,8 @@ def test_admin_can_create_client_token_and_client_can_rotate_it(self, tmp_path, "smart-chat", "private-chat", ) + assert updated_auth.client.policy_actions["pii"] == "audit" + assert updated_auth.client.policy_actions["pci"] == "block" gateway_usage_store(config).record( updated_auth.client, @@ -827,6 +923,61 @@ def test_admin_can_create_client_token_and_client_can_rotate_it(self, tmp_path, assert newest_token != new_token assert authenticate_gateway_request({"authorization": f"Bearer {newest_token}"}, config).allowed + def test_admin_can_update_provider_secret_encrypted(self, tmp_path, monkeypatch): + fastapi_testclient = pytest.importorskip("fastapi.testclient") + monkeypatch.setenv("SENTINELGUARD_ADMIN_PASSWORD", "admin-pass") + monkeypatch.setenv("SENTINELGUARD_VIEWER_PASSWORD", "viewer-pass") + monkeypatch.setenv("SENTINELGUARD_ENCRYPTION_KEY", "test-master-key") + monkeypatch.delenv("OPENAI_API_KEY", raising=False) + state_path = tmp_path / "gateway-provider-secrets.sqlite3" + config = GatewayConfig( + state_backend="sqlite", + state_path=str(state_path), + providers=[ProviderConfig(name="openai-fast", provider="openai", model_name="fast-chat")], + ) + app = create_gateway_app( + guard_config=GuardConfig.preset_minimal(), + gateway_config=config, + ) + client = fastapi_testclient.TestClient(app) + assert client.post( + "/admin/api/login", + json={"username": "admin", "password": "admin-pass"}, + ).status_code == 200 + + raw_provider_key = "sk-test-provider-secret-value" + updated = client.patch( + "/admin/api/providers/openai-fast/secret", + json={"api_key": raw_provider_key}, + ) + + assert updated.status_code == 200 + assert raw_provider_key not in json.dumps(updated.json()) + provider_secret = updated.json()["provider_secret"] + assert provider_secret["secret_hint"] == "configured ...alue" + + with sqlite3.connect(state_path) as connection: + row = connection.execute( + "SELECT encrypted_api_key, secret_hint FROM gateway_provider_secrets WHERE provider_name = ?", + ("openai-fast",), + ).fetchone() + assert row is not None + assert raw_provider_key not in row[0] + assert row[1] == "configured ...alue" + + admin_store = gateway_admin_store(config) + assert admin_store.provider_secret("openai-fast") == raw_provider_key + assert configured_providers(config)[0].api_key == raw_provider_key + + summary = client.get("/admin/api/summary").json() + provider = summary["provider_secrets"]["providers"][0] + assert provider["api_key_source"] == "dashboard" + assert provider["secret_hint"] == "configured ...alue" + + removed = client.delete("/admin/api/providers/openai-fast/secret") + assert removed.status_code == 200 + assert configured_providers(config)[0].api_key is None + def test_viewer_can_read_but_cannot_create_client_token(self, tmp_path, monkeypatch): fastapi_testclient = pytest.importorskip("fastapi.testclient") monkeypatch.setenv("SENTINELGUARD_ADMIN_PASSWORD", "admin-pass")