use std::fmt;
use std::pin::Pin;
use std::sync::Arc;
use async_trait::async_trait;
use futures::Stream;
use futures::StreamExt;
use serde_json::Value;
use tokio_util::sync::CancellationToken as TokioCancellationToken;
use crate::error::ProviderError;
use crate::http;
use crate::key_source::KeySource;
use crate::protocol::{AuthMethod, ProtocolAdapter};
use crate::traits::LlmProvider;
use crate::types::{CompletionRequest, CompletionResponse, ModelInfo, RequestOptions, StreamEvent};
pub struct GenericProvider {
name: String,
adapter: Arc<dyn ProtocolAdapter>,
auth: AuthMethod,
key_source: Option<Arc<dyn KeySource>>,
base_url: String,
models: Vec<ModelInfo>,
client: reqwest::Client,
extra_headers: Vec<(String, String)>,
}
impl fmt::Debug for GenericProvider {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
f.debug_struct("GenericProvider")
.field("name", &self.name)
.field("adapter", &self.adapter)
.field("auth", &self.auth)
.field("key_source", &self.key_source.as_ref().map(|_| ".."))
.field("base_url", &self.base_url)
.field("models", &self.models)
.field("extra_headers", &self.extra_headers)
.finish()
}
}
impl GenericProvider {
pub fn new(
name: String,
adapter: Arc<dyn ProtocolAdapter>,
auth: AuthMethod,
base_url: String,
models: Vec<ModelInfo>,
client: reqwest::Client,
) -> Self {
Self::with_headers_and_key_source(name, adapter, auth, base_url, models, client, Vec::new(), None)
}
pub fn with_headers(
name: String,
adapter: Arc<dyn ProtocolAdapter>,
auth: AuthMethod,
base_url: String,
models: Vec<ModelInfo>,
client: reqwest::Client,
extra_headers: Vec<(String, String)>,
) -> Self {
Self { name, adapter, auth, key_source: None, base_url, models, client, extra_headers }
}
pub fn with_headers_and_key_source(
name: String,
adapter: Arc<dyn ProtocolAdapter>,
auth: AuthMethod,
base_url: String,
models: Vec<ModelInfo>,
client: reqwest::Client,
extra_headers: Vec<(String, String)>,
key_source: Option<Arc<dyn KeySource>>,
) -> Self {
Self { name, adapter, auth, key_source, base_url, models, client, extra_headers }
}
async fn resolve_effective_auth(&self) -> Result<AuthMethod, ProviderError> {
if let Some(ref ks) = self.key_source {
match ks.get_api_key().await {
Ok(key) if !key.is_empty() => {
return Ok(match &self.auth {
AuthMethod::ApiKey { header_name, .. } => {
AuthMethod::ApiKey { header_name: header_name.clone(), key }
}
_ => AuthMethod::Bearer { token: key },
});
}
Ok(_) => {}
Err(e) => return Err(e),
}
}
Ok(self.auth.clone())
}
}
#[async_trait]
impl LlmProvider for GenericProvider {
async fn complete(
&self,
request: CompletionRequest,
options: RequestOptions,
) -> Result<CompletionResponse, ProviderError> {
if let Some(ref ct) = options.cancel {
if ct.is_cancelled() {
return Err(ProviderError::Cancelled);
}
}
let start = std::time::Instant::now();
let url = format!("{}{}", self.base_url, self.adapter.endpoint_path());
let body = self.adapter.build_request_body(&request, false)?;
let effective_auth = self.resolve_effective_auth().await?;
let mut headers = self.adapter.build_auth_headers(&effective_auth);
headers.extend(self.extra_headers.clone());
let resp = http::build_request(&self.client, &url, body, headers, options.timeout).await?;
let status = resp.status();
if !status.is_success() {
let headers = resp.headers().clone();
let error_text = resp.text().await.unwrap_or_default();
return Err(http::handle_error_response(status.as_u16(), &headers, &error_text));
}
let response_text = resp.text().await?;
let response_json: Value =
serde_json::from_str(&response_text).map_err(|e| {
tracing::warn!(
body_preview = %response_text.chars().take(500).collect::<String>(),
body_len = response_text.len(),
"API returned non-JSON response"
);
ProviderError::Format(format!("error decoding response body: {e}"))
})?;
let mut result = self.adapter.parse_response(&response_json)?;
result.latency_ms = start.elapsed().as_millis() as u64;
Ok(result)
}
async fn complete_stream(
&self,
request: CompletionRequest,
options: RequestOptions,
) -> Result<Pin<Box<dyn Stream<Item = Result<StreamEvent, ProviderError>> + Send>>, ProviderError>
{
if let Some(ref ct) = options.cancel {
if ct.is_cancelled() {
return Err(ProviderError::Cancelled);
}
}
let url = format!("{}{}", self.base_url, self.adapter.endpoint_path());
let body = self.adapter.build_request_body(&request, true)?;
let effective_auth = self.resolve_effective_auth().await?;
let mut headers = self.adapter.build_auth_headers(&effective_auth);
headers.extend(self.extra_headers.clone());
let resp = http::build_request(&self.client, &url, body, headers, options.timeout).await?;
let status = resp.status();
if !status.is_success() {
let headers = resp.headers().clone();
let error_text = resp.text().await.unwrap_or_default();
return Err(http::handle_error_response(status.as_u16(), &headers, &error_text));
}
let line_stream = http::create_sse_stream(resp, None);
let adapter = Arc::clone(&self.adapter);
let cancel_opt = options.cancel;
let event_stream = line_stream.filter_map(move |line_result| {
let line = match line_result {
Ok(l) => l,
Err(e) => return futures::future::ready(Some(Err(e))),
};
if line.is_empty() {
return futures::future::ready(None);
}
if let Some(data) = line.strip_prefix("data: ") {
match adapter.parse_sse_event(data) {
Ok(Some(event)) => futures::future::ready(Some(Ok(event))),
Ok(None) => futures::future::ready(None),
Err(e) => futures::future::ready(Some(Err(e))),
}
} else {
futures::future::ready(None)
}
});
let stream: Pin<Box<dyn Stream<Item = Result<StreamEvent, ProviderError>> + Send>> =
match cancel_opt {
Some(ct) => {
let inner: TokioCancellationToken = ct.into();
Box::pin(event_stream.take_until(inner.cancelled_owned()))
}
None => Box::pin(event_stream),
};
Ok(stream)
}
fn models(&self) -> &[ModelInfo] {
&self.models
}
fn name(&self) -> &str {
&self.name
}
}
#[cfg(test)]
mod tests {
use std::sync::Arc;
use std::time::Duration;
use futures::StreamExt;
use serde_json::Value;
use wiremock::matchers::{method, path};
use wiremock::{Mock, MockServer, ResponseTemplate};
use super::*;
use crate::cancel::CancellationToken;
use crate::protocol::ProtocolAdapter;
use crate::types::{
FinishReason, Message, ModelCapabilities, ModelLimits, ModelPricing, StreamEvent,
TokenUsage,
};
#[derive(Debug, Clone)]
struct MockAdapter {
endpoint: String,
proto_name: String,
}
impl ProtocolAdapter for MockAdapter {
fn endpoint_path(&self) -> &str {
&self.endpoint
}
fn build_request_body(
&self,
request: &CompletionRequest,
stream: bool,
) -> Result<Value, ProviderError> {
Ok(serde_json::json!({
"model": request.model.as_deref().unwrap_or("unknown"),
"messages": [],
"stream": stream,
}))
}
fn build_auth_headers(&self, auth: &AuthMethod) -> Vec<(String, String)> {
match auth {
AuthMethod::Bearer { token } => {
vec![("Authorization".to_owned(), format!("Bearer {}", token))]
}
AuthMethod::ApiKey { header_name, key } => {
vec![(header_name.clone(), key.clone())]
}
AuthMethod::None => vec![],
}
}
fn parse_response(&self, body: &Value) -> Result<CompletionResponse, ProviderError> {
let content = body["content"].as_str().map(|s| s.to_owned());
let model = body["model"].as_str().unwrap_or("unknown").to_owned();
Ok(CompletionResponse {
content,
thinking: None,
tool_calls: vec![],
usage: TokenUsage::new(10, 5),
model,
finish_reason: FinishReason::Stop,
latency_ms: 0,
cache_info: None,
..Default::default()
})
}
fn parse_sse_event(&self, data: &str) -> Result<Option<StreamEvent>, ProviderError> {
if data == "[DONE]" {
return Ok(None);
}
let parsed: Value =
serde_json::from_str(data).map_err(|e| ProviderError::Format(e.to_string()))?;
if let Some(delta) = parsed["delta"].as_str() {
Ok(Some(StreamEvent::ContentDelta { delta: delta.to_owned() }))
} else if parsed["done"].as_bool() == Some(true) {
Ok(Some(StreamEvent::Done { finish_reason: FinishReason::Stop, usage: None }))
} else {
Ok(None)
}
}
fn protocol_name(&self) -> &str {
&self.proto_name
}
}
fn make_model(name: &str) -> ModelInfo {
ModelInfo {
name: name.to_owned(),
display_name: None,
provider: None,
capabilities: ModelCapabilities::default(),
pricing: ModelPricing::default(),
limits: ModelLimits::default(),
}
}
fn make_adapter() -> MockAdapter {
MockAdapter { endpoint: "/v1/test".to_owned(), proto_name: "mock".to_owned() }
}
fn make_provider(base_url: &str) -> GenericProvider {
GenericProvider::new(
"test-provider".to_owned(),
Arc::new(make_adapter()),
AuthMethod::None,
base_url.to_owned(),
vec![make_model("test-model"), make_model("test-model-2")],
reqwest::Client::new(),
)
}
fn make_request() -> CompletionRequest {
CompletionRequest::new("test-model", vec![Message::user("Hello")])
}
#[test]
fn test_name_and_models() {
let provider = make_provider("http://localhost");
assert_eq!(provider.name(), "test-provider");
let models = provider.models();
assert_eq!(models.len(), 2);
assert_eq!(models[0].name, "test-model");
assert_eq!(models[1].name, "test-model-2");
}
#[test]
fn test_default_model() {
let provider = make_provider("http://localhost");
assert_eq!(provider.default_model(), Some("test-model"));
}
#[test]
fn test_default_model_empty_models() {
let provider = GenericProvider::new(
"empty".to_owned(),
Arc::new(make_adapter()),
AuthMethod::None,
"http://localhost".to_owned(),
vec![],
reqwest::Client::new(),
);
assert_eq!(provider.default_model(), None);
}
#[test]
fn test_debug_output() {
let provider = make_provider("http://localhost");
let debug_str = format!("{:?}", provider);
assert!(debug_str.contains("GenericProvider"));
assert!(debug_str.contains("test-provider"));
}
#[tokio::test]
async fn test_complete_returns_correct_response() {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/test"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"content": "Hello from mock!",
"model": "test-model",
})))
.mount(&mock_server)
.await;
let provider = make_provider(&mock_server.uri());
let resp = provider.complete(make_request(), RequestOptions::default()).await.unwrap();
assert_eq!(resp.content.as_deref(), Some("Hello from mock!"));
assert_eq!(resp.model, "test-model");
assert_eq!(resp.finish_reason, FinishReason::Stop);
assert_eq!(resp.usage.prompt_tokens, 10);
assert_eq!(resp.usage.completion_tokens, 5);
assert!(resp.latency_ms > 0);
}
#[tokio::test]
async fn test_complete_stream_returns_correct_events() {
let mock_server = MockServer::start().await;
let sse_body = "\
data: {\"delta\":\"Hello\"}\n\
\n\
data: {\"delta\":\" world!\"}\n\
\n\
data: {\"done\":true}\n\
\n\
data: [DONE]\n\n";
Mock::given(method("POST"))
.and(path("/v1/test"))
.respond_with(ResponseTemplate::new(200).set_body_string(sse_body))
.mount(&mock_server)
.await;
let provider = make_provider(&mock_server.uri());
let stream =
provider.complete_stream(make_request(), RequestOptions::default()).await.unwrap();
let events: Vec<StreamEvent> =
stream.filter_map(|r| futures::future::ready(r.ok())).collect().await;
assert_eq!(events.len(), 3, "expected 3 events: 2 deltas + 1 Done, got {events:?}");
assert_eq!(events[0], StreamEvent::ContentDelta { delta: "Hello".to_owned() });
assert_eq!(events[1], StreamEvent::ContentDelta { delta: " world!".to_owned() });
assert_eq!(events[2], StreamEvent::Done { finish_reason: FinishReason::Stop, usage: None });
}
#[tokio::test]
async fn test_timeout_applied_to_complete() {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/test"))
.respond_with(
ResponseTemplate::new(200)
.set_body_json(serde_json::json!({"content": "ok", "model": "m"}))
.set_delay(Duration::from_millis(5000)),
)
.mount(&mock_server)
.await;
let provider = make_provider(&mock_server.uri());
let opts =
RequestOptions { timeout: Some(Duration::from_millis(100)), ..Default::default() };
let result = provider.complete(make_request(), opts).await;
assert!(matches!(result, Err(ProviderError::Timeout { .. })), "expected Timeout");
}
#[tokio::test]
async fn test_timeout_applied_to_complete_stream() {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/test"))
.respond_with(
ResponseTemplate::new(200)
.set_body_string("data: {\"delta\":\"Hi\"}\n\n")
.set_delay(Duration::from_millis(5000)),
)
.mount(&mock_server)
.await;
let provider = make_provider(&mock_server.uri());
let opts =
RequestOptions { timeout: Some(Duration::from_millis(100)), ..Default::default() };
let result = provider.complete_stream(make_request(), opts).await;
assert!(matches!(result, Err(ProviderError::Timeout { .. })), "expected Timeout");
}
#[tokio::test]
async fn test_pre_cancelled_complete_returns_immediate_error() {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/test"))
.respond_with(ResponseTemplate::new(200))
.mount(&mock_server)
.await;
let provider = make_provider(&mock_server.uri());
let cancel = CancellationToken::new();
cancel.cancel();
let opts = RequestOptions { cancel: Some(cancel), ..Default::default() };
let result = provider.complete(make_request(), opts).await;
assert!(matches!(result, Err(ProviderError::Cancelled)), "expected Cancelled");
}
#[tokio::test]
async fn test_pre_cancelled_complete_stream_returns_immediate_error() {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/test"))
.respond_with(ResponseTemplate::new(200))
.mount(&mock_server)
.await;
let provider = make_provider(&mock_server.uri());
let cancel = CancellationToken::new();
cancel.cancel();
let opts = RequestOptions { cancel: Some(cancel), ..Default::default() };
let result = provider.complete_stream(make_request(), opts).await;
assert!(matches!(result, Err(ProviderError::Cancelled)), "expected Cancelled");
}
#[tokio::test]
async fn test_cancellation_stops_streaming() {
let mock_server = MockServer::start().await;
let mut sse_body = String::new();
for i in 0..50 {
sse_body.push_str(&format!("data: {{\"delta\":\"chunk{i}\"}}\n\n"));
}
Mock::given(method("POST"))
.and(path("/v1/test"))
.respond_with(ResponseTemplate::new(200).set_body_string(sse_body))
.mount(&mock_server)
.await;
let provider = make_provider(&mock_server.uri());
let cancel = CancellationToken::new();
let opts = RequestOptions { cancel: Some(cancel.clone()), ..Default::default() };
let stream = provider.complete_stream(make_request(), opts).await.unwrap();
cancel.cancel();
let events: Vec<StreamEvent> =
stream.filter_map(|r| futures::future::ready(r.ok())).collect().await;
assert!(
events.len() < 50,
"expected fewer than 50 events with mid-stream cancel, got {}",
events.len()
);
}
#[tokio::test]
async fn test_http_error_in_complete() {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/test"))
.respond_with(ResponseTemplate::new(500).set_body_string("Internal Server Error"))
.mount(&mock_server)
.await;
let provider = make_provider(&mock_server.uri());
let result = provider.complete(make_request(), RequestOptions::default()).await;
match result {
Err(ProviderError::Internal { status, message }) => {
assert_eq!(status, 500);
assert!(message.contains("Internal Server Error"));
}
other => panic!("expected Internal error, got {other:?}"),
}
}
#[tokio::test]
async fn test_http_error_in_complete_stream() {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/test"))
.respond_with(ResponseTemplate::new(503))
.mount(&mock_server)
.await;
let provider = make_provider(&mock_server.uri());
let result = provider.complete_stream(make_request(), RequestOptions::default()).await;
assert!(matches!(result, Err(ProviderError::Overloaded)), "expected Overloaded");
}
#[tokio::test]
async fn test_complete_401_error() {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/test"))
.respond_with(ResponseTemplate::new(401).set_body_string("Unauthorized"))
.mount(&mock_server)
.await;
let provider = make_provider(&mock_server.uri());
let result = provider.complete(make_request(), RequestOptions::default()).await;
assert!(matches!(result, Err(ProviderError::Auth(_))), "expected Auth, got {result:?}");
}
#[tokio::test]
async fn test_complete_stream_429_error() {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/test"))
.respond_with(
ResponseTemplate::new(429)
.set_body_string("Rate limited")
.insert_header("Retry-After", "30"),
)
.mount(&mock_server)
.await;
let provider = make_provider(&mock_server.uri());
let result = provider.complete_stream(make_request(), RequestOptions::default()).await;
match result {
Err(ProviderError::RateLimit { retry_after_ms }) => {
assert_eq!(retry_after_ms, 30_000);
}
_other => panic!("expected RateLimit, got unexpected variant"),
}
}
#[tokio::test]
async fn test_complete_with_bearer_auth() {
let mock_server = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/test"))
.and(wiremock::matchers::header("Authorization", "Bearer sk-test-key"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"content": "authed!",
"model": "test-model",
})))
.mount(&mock_server)
.await;
let provider = GenericProvider::new(
"auth-provider".to_owned(),
Arc::new(make_adapter()),
AuthMethod::Bearer { token: "sk-test-key".to_owned() },
mock_server.uri(),
vec![make_model("test-model")],
reqwest::Client::new(),
);
let resp = provider.complete(make_request(), RequestOptions::default()).await.unwrap();
assert_eq!(resp.content.as_deref(), Some("authed!"));
}
}