use std::sync::Arc;
use std::time::Duration;
use crate::providers::vertex_auth::VertexAuth;
use crate::resilience::{CircuitBreaker, Profile};
use crate::vertex_endpoint::vertex_endpoint_base;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum ProviderKind {
Vertex,
OpenAiCompat,
}
pub struct Provider {
pub id: &'static str,
pub kind: ProviderKind,
pub client: genai::Client,
pub breaker: Arc<CircuitBreaker>,
pub profile: Profile,
pub label: &'static str,
}
impl std::fmt::Debug for Provider {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("Provider")
.field("id", &self.id)
.field("kind", &self.kind)
.field("label", &self.label)
.field("profile", &self.profile)
.finish_non_exhaustive()
}
}
#[derive(Debug, Clone)]
pub struct VertexProviderConfig {
pub project: String,
pub region: String,
pub request_timeout: Duration,
pub endpoint_override: Option<String>,
}
#[derive(Debug, Clone)]
pub struct OpenAiCompatConfig {
pub base_url: String,
pub api_key: String,
pub request_timeout: Duration,
pub endpoint_override: Option<String>,
}
pub fn build_vertex_provider(
id: &'static str,
config: VertexProviderConfig,
auth: Arc<VertexAuth>,
) -> anyhow::Result<Provider> {
let endpoint_base = config
.endpoint_override
.clone()
.unwrap_or_else(|| vertex_endpoint_base(&config.region));
let project = config.project.clone();
let region = config.region.clone();
let client = genai::Client::builder()
.with_auth_resolver(auth.into_auth_resolver())
.with_service_target_resolver_fn(
move |service_target: genai::ServiceTarget| -> genai::resolver::Result<genai::ServiceTarget> {
let endpoint = genai::resolver::Endpoint::from_owned(format!(
"{endpoint_base}/v1/projects/{project}/locations/{region}/",
));
Ok(genai::ServiceTarget {
endpoint,
auth: service_target.auth,
model: service_target.model,
})
},
)
.with_adapter_kind(genai::adapter::AdapterKind::Vertex)
.with_web_config(genai::WebConfig::default().with_timeout(config.request_timeout))
.build();
let breaker = Arc::new(CircuitBreaker::new(id, Profile::Aggressive));
Ok(Provider {
id,
kind: ProviderKind::Vertex,
client,
breaker,
profile: Profile::Aggressive,
label: id,
})
}
fn derive_openai_compat_endpoint(base_url: &str) -> String {
let mut s = base_url.trim_end_matches('/').to_string();
s.push('/');
s
}
pub fn build_openai_compat_provider(
id: &'static str,
config: OpenAiCompatConfig,
) -> anyhow::Result<Provider> {
let api_key = config.api_key.clone();
let endpoint_url = config
.endpoint_override
.clone()
.unwrap_or_else(|| derive_openai_compat_endpoint(&config.base_url));
let client = genai::Client::builder()
.with_auth_resolver(genai::resolver::AuthResolver::from_resolver_fn(
move |_model_iden: genai::ModelIden| -> genai::resolver::Result<Option<genai::resolver::AuthData>> {
Ok(Some(genai::resolver::AuthData::from_single(api_key.clone())))
},
))
.with_service_target_resolver_fn(
move |service_target: genai::ServiceTarget| -> genai::resolver::Result<genai::ServiceTarget> {
let endpoint = genai::resolver::Endpoint::from_owned(endpoint_url.clone());
Ok(genai::ServiceTarget {
endpoint,
auth: service_target.auth,
model: service_target.model,
})
},
)
.with_adapter_kind(genai::adapter::AdapterKind::OpenAI)
.with_web_config(genai::WebConfig::default().with_timeout(config.request_timeout))
.build();
let breaker = Arc::new(CircuitBreaker::new(id, Profile::Default));
Ok(Provider {
id,
kind: ProviderKind::OpenAiCompat,
client,
breaker,
profile: Profile::Default,
label: id,
})
}
#[cfg(test)]
mod tests {
use super::*;
use std::time::Duration;
fn dummy_vertex_auth() -> Arc<VertexAuth> {
Arc::new(VertexAuth::with_fetcher(|| {
Box::pin(async { Ok(("test-token".into(), Duration::from_secs(3600))) })
}))
}
#[test]
fn build_vertex_provider_yields_vertex_kind_and_aggressive_profile() {
let p = build_vertex_provider(
"vertex_global",
VertexProviderConfig {
project: "test-proj".into(),
region: "global".into(),
request_timeout: Duration::from_secs(120),
endpoint_override: None,
},
dummy_vertex_auth(),
)
.unwrap();
assert_eq!(p.id, "vertex_global");
assert_eq!(p.kind, ProviderKind::Vertex);
assert!(matches!(p.profile, Profile::Aggressive));
}
#[test]
fn build_openai_compat_provider_yields_oai_kind_and_default_profile() {
let p = build_openai_compat_provider(
"qwen",
OpenAiCompatConfig {
base_url: "https://dashscope-intl.aliyuncs.com/compatible-mode/v1".into(),
api_key: "sk-test".into(),
request_timeout: Duration::from_secs(120),
endpoint_override: None,
},
)
.unwrap();
assert_eq!(p.id, "qwen");
assert_eq!(p.kind, ProviderKind::OpenAiCompat);
assert!(matches!(p.profile, Profile::Default));
}
#[tokio::test]
async fn vertex_provider_sends_bearer_token_on_request() {
use wiremock::matchers::{header, method, path_regex};
use wiremock::{Mock, MockServer, Request, ResponseTemplate};
let mock = MockServer::start().await;
Mock::given(method("POST"))
.and(path_regex(
r"^/v1/projects/test-proj/locations/global/publishers/google/models/gemini-2\.5-flash:generateContent$",
))
.and(header("authorization", "Bearer test-token"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"candidates": [{
"content": { "parts": [{ "text": "ok" }], "role": "model" },
"finishReason": "STOP"
}],
"modelVersion": "gemini-2.5-flash",
"usageMetadata": { "promptTokenCount": 1, "candidatesTokenCount": 1, "totalTokenCount": 2 }
})))
.mount(&mock)
.await;
let provider = build_vertex_provider(
"vertex_global",
VertexProviderConfig {
project: "test-proj".into(),
region: "global".into(),
request_timeout: Duration::from_secs(5),
endpoint_override: Some(mock.uri()),
},
dummy_vertex_auth(),
)
.unwrap();
let _ = provider
.client
.exec_chat(
"gemini-2.5-flash",
genai::chat::ChatRequest::from_user("hi"),
None,
)
.await;
let received: Vec<Request> = mock.received_requests().await.unwrap_or_default();
assert!(
received.iter().any(|r| {
r.headers
.get("authorization")
.and_then(|v| v.to_str().ok())
== Some("Bearer test-token")
}),
"expected Authorization: Bearer test-token on at least one outbound request, got headers: {:?}",
received.iter().map(|r| r.headers.clone()).collect::<Vec<_>>(),
);
assert!(
received.iter().any(|r| {
let p = r.url.path();
p == "/v1/projects/test-proj/locations/global/publishers/google/models/gemini-2.5-flash:generateContent"
}),
"expected exact Vertex URL with no literal {{model}} and no doubled :generateContent, got URLs: {:?}",
received.iter().map(|r| r.url.path().to_string()).collect::<Vec<_>>(),
);
}
#[tokio::test]
async fn openai_compat_provider_sends_bearer_token_on_request() {
use wiremock::matchers::{header, method, path};
use wiremock::{Mock, MockServer, Request, ResponseTemplate};
let mock = MockServer::start().await;
Mock::given(method("POST"))
.and(path("/v1/chat/completions"))
.and(header("authorization", "Bearer test-key"))
.respond_with(ResponseTemplate::new(200).set_body_json(serde_json::json!({
"id": "chatcmpl-test",
"object": "chat.completion",
"created": 0,
"model": "qwen-vl-max",
"choices": [{
"index": 0,
"message": { "role": "assistant", "content": "ok" },
"finish_reason": "stop"
}],
"usage": { "prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2 }
})))
.mount(&mock)
.await;
let provider = build_openai_compat_provider(
"qwen",
OpenAiCompatConfig {
base_url: format!("{}/v1", mock.uri()),
api_key: "test-key".into(),
request_timeout: Duration::from_secs(5),
endpoint_override: None,
},
)
.unwrap();
let _ = provider
.client
.exec_chat(
"qwen-vl-max",
genai::chat::ChatRequest::from_user("hi"),
None,
)
.await;
let received: Vec<Request> = mock.received_requests().await.unwrap_or_default();
assert!(
received.iter().any(|r| {
r.headers
.get("authorization")
.and_then(|v| v.to_str().ok())
== Some("Bearer test-key")
}),
"expected Authorization: Bearer test-key on at least one outbound request, got headers: {:?}",
received.iter().map(|r| r.headers.clone()).collect::<Vec<_>>(),
);
assert!(
received
.iter()
.any(|r| r.url.path() == "/v1/chat/completions"),
"expected exact path /v1/chat/completions, got: {:?}",
received
.iter()
.map(|r| r.url.path().to_string())
.collect::<Vec<_>>(),
);
}
#[test]
fn openai_compat_endpoint_resolves_to_chat_completions() {
use reqwest::Url;
let base = "https://dashscope-intl.aliyuncs.com/compatible-mode/v1";
let endpoint = derive_openai_compat_endpoint(base);
let parsed = Url::parse(&endpoint).unwrap();
let joined = parsed.join("chat/completions").unwrap();
assert_eq!(
joined.as_str(),
"https://dashscope-intl.aliyuncs.com/compatible-mode/v1/chat/completions",
);
let _ = build_openai_compat_provider(
"qwen",
OpenAiCompatConfig {
base_url: base.into(),
api_key: "k".into(),
request_timeout: Duration::from_secs(5),
endpoint_override: None,
},
)
.unwrap();
}
}