use std::{
collections::{BTreeMap, BTreeSet},
sync::{Arc, Mutex},
};
use serde::{Deserialize, Serialize};
use serde_json::{Map, Value};
use super::{HttpMethod, HttpRequest};
pub type DynProviderRequestAuditRecorder = Arc<dyn ProviderRequestAuditRecorder>;
#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum ProviderRequestAuditPayloadPolicy {
#[default]
Omit,
Redacted,
Full,
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct ProviderRequestAuditPolicy {
#[serde(default)]
pub headers: ProviderRequestAuditPayloadPolicy,
#[serde(default)]
pub body: ProviderRequestAuditPayloadPolicy,
#[serde(default = "default_sensitive_header_keys")]
pub sensitive_header_keys: BTreeSet<String>,
#[serde(default = "default_sensitive_body_keys")]
pub sensitive_body_keys: BTreeSet<String>,
#[serde(default = "default_redaction_value")]
pub redaction_value: Value,
}
impl Default for ProviderRequestAuditPolicy {
fn default() -> Self {
Self {
headers: ProviderRequestAuditPayloadPolicy::Omit,
body: ProviderRequestAuditPayloadPolicy::Omit,
sensitive_header_keys: default_sensitive_header_keys(),
sensitive_body_keys: default_sensitive_body_keys(),
redaction_value: default_redaction_value(),
}
}
}
impl ProviderRequestAuditPolicy {
#[must_use]
pub fn metadata_only() -> Self {
Self::default()
}
#[must_use]
pub fn redacted_payloads() -> Self {
Self {
headers: ProviderRequestAuditPayloadPolicy::Redacted,
body: ProviderRequestAuditPayloadPolicy::Redacted,
..Self::default()
}
}
#[must_use]
pub fn full_payloads() -> Self {
Self {
headers: ProviderRequestAuditPayloadPolicy::Full,
body: ProviderRequestAuditPayloadPolicy::Full,
..Self::default()
}
}
#[must_use]
pub fn with_sensitive_header_key(mut self, key: impl AsRef<str>) -> Self {
self.sensitive_header_keys
.insert(normalize_key(key.as_ref()));
self
}
#[must_use]
pub fn with_sensitive_body_key(mut self, key: impl AsRef<str>) -> Self {
self.sensitive_body_keys.insert(normalize_key(key.as_ref()));
self
}
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct ProviderRequestAuditSnapshot {
pub provider_name: String,
pub model_name: String,
pub stream: bool,
pub method: HttpMethod,
pub url: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub headers: Option<BTreeMap<String, String>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub body: Option<Value>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub timeout_ms: Option<u64>,
#[serde(default, skip_serializing_if = "Map::is_empty")]
pub metadata: Map<String, Value>,
}
impl ProviderRequestAuditSnapshot {
#[must_use]
pub fn from_request(
provider_name: impl Into<String>,
model_name: impl Into<String>,
stream: bool,
request: &HttpRequest,
policy: &ProviderRequestAuditPolicy,
) -> Self {
Self {
provider_name: provider_name.into(),
model_name: model_name.into(),
stream,
method: request.method,
url: request.url.clone(),
headers: capture_headers(&request.headers, policy),
body: capture_body(&request.body, policy),
timeout_ms: request
.timeout
.and_then(|duration| u64::try_from(duration.as_millis()).ok()),
metadata: request.metadata.clone(),
}
}
}
pub trait ProviderRequestAuditRecorder: Send + Sync {
fn record_provider_request(&self, snapshot: ProviderRequestAuditSnapshot);
}
#[derive(Default)]
pub struct InMemoryProviderRequestAuditRecorder {
snapshots: Mutex<Vec<ProviderRequestAuditSnapshot>>,
}
impl InMemoryProviderRequestAuditRecorder {
#[must_use]
pub fn new() -> Self {
Self::default()
}
#[must_use]
pub fn snapshots(&self) -> Vec<ProviderRequestAuditSnapshot> {
self.snapshots
.lock()
.map_or_else(|_| Vec::new(), |v| v.clone())
}
#[must_use]
pub fn take_snapshots(&self) -> Vec<ProviderRequestAuditSnapshot> {
self.snapshots
.lock()
.map_or_else(|_| Vec::new(), |mut v| std::mem::take(&mut *v))
}
}
impl ProviderRequestAuditRecorder for InMemoryProviderRequestAuditRecorder {
fn record_provider_request(&self, snapshot: ProviderRequestAuditSnapshot) {
if let Ok(mut snapshots) = self.snapshots.lock() {
snapshots.push(snapshot);
}
}
}
#[derive(Clone)]
pub struct ProviderRequestAuditCapture {
recorder: DynProviderRequestAuditRecorder,
policy: ProviderRequestAuditPolicy,
}
impl ProviderRequestAuditCapture {
pub(crate) fn new(
recorder: DynProviderRequestAuditRecorder,
policy: ProviderRequestAuditPolicy,
) -> Self {
Self { recorder, policy }
}
pub(crate) fn record(
&self,
provider_name: &str,
model_name: &str,
stream: bool,
request: &HttpRequest,
) {
self.recorder
.record_provider_request(ProviderRequestAuditSnapshot::from_request(
provider_name,
model_name,
stream,
request,
&self.policy,
));
}
}
fn capture_headers(
headers: &BTreeMap<String, String>,
policy: &ProviderRequestAuditPolicy,
) -> Option<BTreeMap<String, String>> {
match policy.headers {
ProviderRequestAuditPayloadPolicy::Omit => None,
ProviderRequestAuditPayloadPolicy::Full => Some(headers.clone()),
ProviderRequestAuditPayloadPolicy::Redacted => Some(
headers
.iter()
.map(|(key, value)| {
if policy.sensitive_header_keys.contains(&normalize_key(key)) {
(key.clone(), policy.redaction_value.to_string())
} else {
(key.clone(), value.clone())
}
})
.collect(),
),
}
}
fn capture_body(body: &Value, policy: &ProviderRequestAuditPolicy) -> Option<Value> {
match policy.body {
ProviderRequestAuditPayloadPolicy::Omit => None,
ProviderRequestAuditPayloadPolicy::Full => Some(body.clone()),
ProviderRequestAuditPayloadPolicy::Redacted => {
let mut body = body.clone();
redact_json_value(&mut body, policy);
Some(body)
}
}
}
fn redact_json_value(value: &mut Value, policy: &ProviderRequestAuditPolicy) {
match value {
Value::Object(object) => {
for (key, value) in object {
if policy.sensitive_body_keys.contains(&normalize_key(key)) {
*value = policy.redaction_value.clone();
} else {
redact_json_value(value, policy);
}
}
}
Value::Array(values) => {
for value in values {
redact_json_value(value, policy);
}
}
Value::Null | Value::Bool(_) | Value::Number(_) | Value::String(_) => {}
}
}
fn normalize_key(key: &str) -> String {
key.to_ascii_lowercase()
}
fn default_sensitive_header_keys() -> BTreeSet<String> {
[
"authorization",
"proxy-authorization",
"x-api-key",
"api-key",
"cookie",
"set-cookie",
]
.into_iter()
.map(str::to_string)
.collect()
}
fn default_sensitive_body_keys() -> BTreeSet<String> {
[
"api_key",
"apikey",
"authorization",
"access_token",
"refresh_token",
"id_token",
"token",
"password",
"secret",
]
.into_iter()
.map(str::to_string)
.collect()
}
fn default_redaction_value() -> Value {
serde_json::json!({ "redacted": true })
}