use std::sync::Arc;
use std::time::Duration;
use crate::ZaiResult;
use crate::client::endpoint::EndpointConfig;
use crate::client::secret::ApiSecret;
#[derive(Clone)]
pub struct ZaiClient {
inner: Arc<ClientInner>,
}
pub(crate) struct ClientInner {
pub(crate) secret: ApiSecret,
pub(crate) endpoints: EndpointConfig,
pub(crate) transport: HttpTransportConfig,
#[allow(dead_code)]
pub(crate) reqwest: reqwest::Client,
pub(crate) sender: crate::client::transport::Transport,
}
impl ZaiClient {
pub fn builder(api_key: impl Into<String>) -> ZaiClientBuilder {
ZaiClientBuilder {
api_key: api_key.into(),
endpoints: EndpointConfig::builder(),
transport: HttpTransportConfig::default(),
allow_insecure: false,
}
}
pub fn from_env() -> ZaiResult<Self> {
let key = std::env::var("ZHIPU_API_KEY").map_err(|_| crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: "ZHIPU_API_KEY environment variable not set".to_string(),
})?;
if key.trim().is_empty() {
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: "ZHIPU_API_KEY environment variable is empty".to_string(),
});
}
Self::builder(key).build()
}
pub fn endpoints(&self) -> &EndpointConfig {
&self.inner.endpoints
}
pub fn transport(&self) -> &HttpTransportConfig {
&self.inner.transport
}
#[allow(dead_code)]
pub(crate) fn reqwest(&self) -> &reqwest::Client {
&self.inner.reqwest
}
pub(crate) fn secret(&self) -> &ApiSecret {
&self.inner.secret
}
pub(crate) async fn send_json<B, R>(
&self,
method: &'static str,
url: String,
body: &B,
) -> ZaiResult<R>
where
B: serde::Serialize + ?Sized,
R: serde::de::DeserializeOwned,
{
let bytes = bytes::Bytes::from(serde_json::to_vec(body).map_err(crate::ZaiError::from)?);
let request = crate::client::transport::request::PreparedRequest {
method,
url,
body: crate::client::transport::request::BodyKind::Bytes(&bytes),
retry_safety: crate::client::transport::retry::RetrySafety::for_method(method),
retry_override: None,
response_mode: crate::client::transport::request::ResponseMode::Json,
route_template: "typed-api-request",
};
self.inner.sender.send(&request).await?.json()
}
pub(crate) async fn send_empty<R>(&self, method: &'static str, url: String) -> ZaiResult<R>
where
R: serde::de::DeserializeOwned,
{
let request = crate::client::transport::request::PreparedRequest {
method,
url,
body: crate::client::transport::request::BodyKind::None,
retry_safety: crate::client::transport::retry::RetrySafety::for_method(method),
retry_override: None,
response_mode: crate::client::transport::request::ResponseMode::Json,
route_template: "typed-api-request",
};
self.inner.sender.send(&request).await?.json()
}
pub(crate) async fn send_json_bytes<B: serde::Serialize + ?Sized>(
&self,
method: &'static str,
url: String,
body: &B,
) -> ZaiResult<bytes::Bytes> {
let bytes = bytes::Bytes::from(serde_json::to_vec(body).map_err(crate::ZaiError::from)?);
let request = crate::client::transport::request::PreparedRequest {
method,
url,
body: crate::client::transport::request::BodyKind::Bytes(&bytes),
retry_safety: crate::client::transport::retry::RetrySafety::for_method(method),
retry_override: None,
response_mode: crate::client::transport::request::ResponseMode::Binary,
route_template: "binary-api-request",
};
self.inner.sender.send(&request).await?.bytes()
}
pub(crate) async fn send_empty_bytes(
&self,
method: &'static str,
url: String,
) -> ZaiResult<bytes::Bytes> {
let request = crate::client::transport::request::PreparedRequest {
method,
url,
body: crate::client::transport::request::BodyKind::None,
retry_safety: crate::client::transport::retry::RetrySafety::for_method(method),
retry_override: None,
response_mode: crate::client::transport::request::ResponseMode::Binary,
route_template: "binary-api-request",
};
self.inner.sender.send(&request).await?.bytes()
}
pub(crate) async fn send_multipart<R: serde::de::DeserializeOwned>(
&self,
method: &'static str,
url: String,
factory: &crate::client::transport::multipart::MultipartBodyFactory,
) -> ZaiResult<R> {
let request = crate::client::transport::request::PreparedRequest {
method,
url,
body: crate::client::transport::request::BodyKind::Multipart(factory),
retry_safety: crate::client::transport::retry::RetrySafety::for_method(method),
retry_override: None,
response_mode: crate::client::transport::request::ResponseMode::Json,
route_template: "multipart-api-request",
};
self.inner.sender.send(&request).await?.json()
}
}
impl std::fmt::Debug for ZaiClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ZaiClient")
.field("secret", &self.inner.secret)
.field("endpoints", &self.inner.endpoints)
.field("transport", &self.inner.transport)
.finish_non_exhaustive()
}
}
pub struct ZaiClientBuilder {
api_key: String,
endpoints: crate::client::endpoint::EndpointConfigBuilder,
transport: HttpTransportConfig,
allow_insecure: bool,
}
impl ZaiClientBuilder {
pub fn endpoint(
mut self,
family: crate::client::endpoint::ApiFamily,
base: &'static str,
) -> Self {
use crate::client::endpoint::ApiFamily::*;
match family {
PaasV4 => self.endpoints = self.endpoints.paas_v4(base),
CodingPaasV4 => self.endpoints = self.endpoints.coding_paas_v4(base),
AgentV1 => self.endpoints = self.endpoints.agent_v1(base),
LlmApplication | ApplicationV2 | ApplicationV3 => {
self.endpoints = self.endpoints.llm_application(base)
},
Zrag => self.endpoints = self.endpoints.zrag(base),
Monitor => self.endpoints = self.endpoints.monitor(base),
Realtime => self.endpoints = self.endpoints.realtime(base),
}
self
}
pub fn transport(mut self, transport: HttpTransportConfig) -> Self {
self.transport = transport;
self
}
pub fn allow_insecure_transport(mut self, allow: bool) -> Self {
self.allow_insecure = allow;
self
}
pub fn build(self) -> ZaiResult<ZaiClient> {
if self.api_key.trim().is_empty() {
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: "ZaiClient requires a non-empty api_key".to_string(),
});
}
let endpoints = self.endpoints.build(self.allow_insecure)?;
let reqwest = build_reqwest_client(&self.transport)?;
let secret = ApiSecret::new(self.api_key);
let sender = crate::client::transport::Transport::new(
reqwest.clone(),
secret.clone(),
&self.transport,
);
let inner = Arc::new(ClientInner {
secret,
endpoints,
transport: self.transport,
reqwest,
sender,
});
Ok(ZaiClient { inner })
}
}
fn build_reqwest_client(transport: &HttpTransportConfig) -> ZaiResult<reqwest::Client> {
let mut builder = reqwest::Client::builder()
.redirect(reqwest::redirect::Policy::none())
.pool_max_idle_per_host(8)
.pool_idle_timeout(Some(Duration::from_secs(90)))
.tcp_keepalive(Some(Duration::from_secs(60)))
.connect_timeout(transport.connect_timeout)
.timeout(transport.request_timeout);
if transport.enable_compression {
builder = builder.gzip(true);
}
builder.build().map_err(crate::ZaiError::from)
}
const ALLOWED_HEADER_NAMES: &[&str] = &["Accept-Language", "X-Correlation-ID", "X-Test-Client"];
#[derive(Debug, Clone)]
pub struct AdditionalHeader {
name: &'static str,
value: String,
}
impl AdditionalHeader {
pub fn new(name: &str, value: &str) -> ZaiResult<Self> {
let allowed = ALLOWED_HEADER_NAMES.contains(&name);
if !allowed {
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: format!("header name {name:?} is not allow-listed"),
});
}
if value.len() > 1024 {
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: "additional header value exceeds 1024 bytes".to_string(),
});
}
if !value.chars().all(|c| !c.is_control()) {
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: "additional header value must be printable".to_string(),
});
}
let static_name = ALLOWED_HEADER_NAMES
.iter()
.copied()
.find(|n| *n == name)
.expect("validated above");
Ok(Self {
name: static_name,
value: value.to_string(),
})
}
pub fn name(&self) -> &'static str {
self.name
}
pub fn value(&self) -> &str {
&self.value
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RetryOverride {
AssumeIdempotent,
}
#[derive(Debug, Clone)]
pub struct HttpTransportConfig {
pub connect_timeout: Duration,
pub request_timeout: Duration,
pub enable_compression: bool,
pub max_attempts: u8,
pub additional_headers: Vec<AdditionalHeader>,
}
impl Default for HttpTransportConfig {
fn default() -> Self {
Self {
connect_timeout: Duration::from_secs(10),
request_timeout: Duration::from_secs(60),
enable_compression: true,
max_attempts: 3,
additional_headers: Vec::new(),
}
}
}
impl HttpTransportConfig {
pub fn builder() -> HttpTransportConfigBuilder {
HttpTransportConfigBuilder {
config: Self::default(),
}
}
pub fn with_request_timeout(mut self, d: Duration) -> ZaiResult<Self> {
if d > Duration::from_secs(60) {
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: "request_timeout may only be lowered (max 60s)".to_string(),
});
}
self.request_timeout = d;
Ok(self)
}
pub fn with_connect_timeout(mut self, d: Duration) -> ZaiResult<Self> {
if d > Duration::from_secs(10) {
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: "connect_timeout may only be lowered (max 10s)".to_string(),
});
}
self.connect_timeout = d;
Ok(self)
}
pub fn with_max_attempts(mut self, n: u8) -> ZaiResult<Self> {
if n == 0 || n > 3 {
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: "max_attempts must be 1, 2 or 3".to_string(),
});
}
self.max_attempts = n;
Ok(self)
}
pub fn with_additional_header(mut self, header: AdditionalHeader) -> Self {
self.additional_headers.push(header);
self
}
}
pub struct HttpTransportConfigBuilder {
config: HttpTransportConfig,
}
impl HttpTransportConfigBuilder {
pub fn request_timeout(mut self, d: Duration) -> ZaiResult<Self> {
self.config = self.config.with_request_timeout(d)?;
Ok(self)
}
pub fn connect_timeout(mut self, d: Duration) -> ZaiResult<Self> {
self.config = self.config.with_connect_timeout(d)?;
Ok(self)
}
pub fn max_attempts(mut self, n: u8) -> ZaiResult<Self> {
self.config = self.config.with_max_attempts(n)?;
Ok(self)
}
pub fn additional_header(mut self, header: AdditionalHeader) -> Self {
self.config.additional_headers.push(header);
self
}
pub fn build(self) -> HttpTransportConfig {
self.config
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn builder_rejects_blank_key() {
assert!(ZaiClient::builder(" ").build().is_err());
assert!(ZaiClient::builder("").build().is_err());
}
#[test]
fn clone_shares_inner_no_secret_leak() {
let c = ZaiClient::builder("abcdefghij.0123456789abcdef")
.build()
.unwrap();
let c2 = c.clone();
let dbg = format!("{c2:?}");
assert!(dbg.contains("[REDACTED]"));
assert!(!dbg.contains("abcdefghij"));
}
#[test]
fn additional_header_allow_list() {
assert!(AdditionalHeader::new("X-Test-Client", "preserved").is_ok());
assert!(AdditionalHeader::new("Authorization", "nope").is_err());
assert!(AdditionalHeader::new("Cookie", "nope").is_err());
assert!(AdditionalHeader::new("Proxy-Authorization", "nope").is_err());
}
#[test]
fn additional_header_value_limits() {
let long = "x".repeat(1025);
assert!(AdditionalHeader::new("X-Test-Client", &long).is_err());
assert!(AdditionalHeader::new("X-Test-Client", "ok\0bad").is_err());
}
#[test]
fn transport_only_tightens() {
assert!(
HttpTransportConfig::default()
.with_request_timeout(Duration::from_secs(120))
.is_err()
);
assert!(
HttpTransportConfig::default()
.with_request_timeout(Duration::from_secs(5))
.is_ok()
);
assert!(HttpTransportConfig::default().with_max_attempts(0).is_err());
assert!(HttpTransportConfig::default().with_max_attempts(4).is_err());
assert!(HttpTransportConfig::default().with_max_attempts(2).is_ok());
}
}