use async_trait::async_trait;
use serde_json::Value;
use std::collections::HashMap;
use super::capabilities::{Capabilities, ProviderInfo};
use super::error::LlmError;
use super::http_client::HttpClient;
use super::request::ChatRequest;
use super::response::{ChatResponse, ChatStream};
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum CallMode {
Stream,
Once,
}
#[derive(Clone)]
pub struct RawRequest {
pub url: String,
pub method: HttpMethod,
pub headers: HashMap<String, String>,
pub body: Value,
pub stream: bool,
}
fn is_sensitive_header(name: &str) -> bool {
let name = name.to_ascii_lowercase();
name.contains("authorization")
|| name.contains("api-key")
|| name.contains("apikey")
|| name.contains("x-api-key")
|| name.contains("token")
}
impl std::fmt::Debug for RawRequest {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
let headers: HashMap<&str, &str> = self
.headers
.iter()
.map(|(k, v)| {
(
k.as_str(),
if is_sensitive_header(k) {
"***"
} else {
v.as_str()
},
)
})
.collect();
f.debug_struct("RawRequest")
.field("url", &self.url)
.field("method", &self.method)
.field("headers", &headers)
.field("body", &self.body)
.field("stream", &self.stream)
.finish()
}
}
#[cfg(test)]
mod debug_tests {
use super::*;
fn request_with(headers: &[(&str, &str)]) -> RawRequest {
RawRequest {
url: "https://api.example.com/v1/messages".to_string(),
method: HttpMethod::Post,
headers: headers
.iter()
.map(|(k, v)| (k.to_string(), v.to_string()))
.collect(),
body: serde_json::json!({"model": "m"}),
stream: true,
}
}
#[test]
fn credential_headers_are_redacted() {
let secret = "sk-super-secret-value";
for name in [
"authorization",
"Authorization",
"x-api-key",
"X-API-KEY",
"api-key",
"apikey",
"x-goog-api-key",
"x-session-token",
"bearer-token",
] {
let rendered = format!("{:?}", request_with(&[(name, secret)]));
assert!(
!rendered.contains(secret),
"header '{name}' leaked its value: {rendered}"
);
assert!(
rendered.contains("***"),
"header '{name}' should be redacted: {rendered}"
);
}
}
#[test]
fn non_credential_headers_stay_visible() {
let rendered = format!(
"{:?}",
request_with(&[
("content-type", "application/json"),
("anthropic-version", "2023-06-01"),
])
);
assert!(rendered.contains("application/json"));
assert!(rendered.contains("2023-06-01"));
assert!(!rendered.contains("***"));
}
#[test]
fn other_fields_remain_visible() {
let rendered = format!("{:?}", request_with(&[("x-api-key", "secret")]));
assert!(rendered.contains("https://api.example.com/v1/messages"));
assert!(rendered.contains("Post"));
assert!(rendered.contains("stream: true"));
assert!(rendered.contains("model"));
}
#[test]
fn is_sensitive_header_recognises_credential_names() {
for name in [
"authorization",
"X-API-Key",
"apikey",
"refresh_token",
"API-KEY",
] {
assert!(is_sensitive_header(name), "{name} should be sensitive");
}
for name in ["content-type", "user-agent", "anthropic-version", "accept"] {
assert!(!is_sensitive_header(name), "{name} should not be sensitive");
}
}
}
#[derive(Debug, Clone, Default, PartialEq, Eq)]
pub enum HttpMethod {
#[default]
Post,
Get,
Put,
Delete,
}
#[derive(Debug, Default)]
#[allow(dead_code)]
pub struct StreamState {
pub data: HashMap<String, Value>,
}
#[async_trait]
pub trait RawAdapter: Send + Sync {
fn build_request(&self, request: &ChatRequest, mode: CallMode) -> Result<RawRequest, LlmError>;
async fn execute_stream(
&self,
client: &dyn HttpClient,
request: RawRequest,
) -> Result<ChatStream, LlmError>;
async fn parse_sse_stream(
&self,
_client: &dyn HttpClient,
_request: RawRequest,
_response: super::http_client::HttpResponse,
) -> Result<ChatStream, LlmError> {
Err(LlmError::llm("parse_sse_stream not implemented"))
}
fn parse_response(&self, _body: &[u8]) -> Result<ChatResponse, LlmError> {
Err(LlmError::llm("Non-streaming mode not supported"))
}
fn capabilities(&self) -> Capabilities;
fn info(&self) -> ProviderInfo;
fn supported_modes(&self) -> &[CallMode] {
&[CallMode::Stream, CallMode::Once]
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn call_mode_equality() {
assert_eq!(CallMode::Stream, CallMode::Stream);
assert_eq!(CallMode::Once, CallMode::Once);
assert_ne!(CallMode::Stream, CallMode::Once);
}
#[test]
fn http_method_default_is_post() {
assert_eq!(HttpMethod::default(), HttpMethod::Post);
}
#[test]
fn raw_request_clone() {
let req = RawRequest {
url: "https://api.example.com/v1/chat".to_string(),
method: HttpMethod::Post,
headers: HashMap::from([("Authorization".to_string(), "Bearer sk-xxx".to_string())]),
body: serde_json::json!({"model": "test"}),
stream: true,
};
let cloned = req.clone();
assert_eq!(cloned.url, req.url);
assert_eq!(cloned.method, req.method);
assert_eq!(cloned.stream, req.stream);
}
#[test]
fn stream_state_default() {
let state = StreamState::default();
assert!(state.data.is_empty());
}
}