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) endpoints: EndpointConfig,
pub(crate) transport: HttpTransportConfig,
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
}
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::Audio,
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::File,
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()
}
pub(crate) async fn send_sse_json<B: serde::Serialize + ?Sized>(
&self,
method: &'static str,
url: String,
body: &B,
) -> ZaiResult<crate::client::transport::SseByteStream> {
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::NonIdempotent,
retry_override: None,
response_mode: crate::client::transport::request::ResponseMode::Json,
route_template: "sse-api-request",
};
self.inner.sender.send_sse(&request).await
}
pub(crate) async fn send_sse_multipart(
&self,
method: &'static str,
url: String,
factory: &crate::client::transport::multipart::MultipartBodyFactory,
) -> ZaiResult<crate::client::transport::SseByteStream> {
let request = crate::client::transport::request::PreparedRequest {
method,
url,
body: crate::client::transport::request::BodyKind::Multipart(factory),
retry_safety: crate::client::transport::retry::RetrySafety::NonIdempotent,
retry_override: None,
response_mode: crate::client::transport::request::ResponseMode::Json,
route_template: "multipart-sse-api-request",
};
self.inner.sender.send_sse(&request).await
}
}
impl std::fmt::Debug for ZaiClient {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("ZaiClient")
.field("credentials", &"[REDACTED]")
.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: impl Into<String>,
) -> 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(),
});
}
if self.api_key.trim() != self.api_key {
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: "ZaiClient api_key must not contain surrounding whitespace".to_string(),
});
}
if !self.api_key.bytes().all(|byte| byte.is_ascii_graphic()) {
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: "ZaiClient api_key must contain printable ASCII without whitespace"
.to_string(),
});
}
self.transport.validate()?;
let endpoints = self.endpoints.build(self.allow_insecure)?;
let reqwest = build_reqwest_client(&self.transport)?;
let sender = crate::client::transport::Transport::new(
reqwest,
ApiSecret::new(self.api_key),
&self.transport,
);
let inner = Arc::new(ClientInner {
endpoints,
transport: self.transport,
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);
builder = builder.gzip(transport.enable_compression);
builder.build().map_err(crate::ZaiError::from)
}
const ALLOWED_HEADER_NAMES: &[&str] = &["Accept-Language", "X-Correlation-ID", "X-Test-Client"];
#[derive(Clone)]
pub struct AdditionalHeader {
name: &'static str,
value: String,
}
impl std::fmt::Debug for AdditionalHeader {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("AdditionalHeader")
.field("name", &self.name)
.field("value", &"[REDACTED]")
.finish()
}
}
impl AdditionalHeader {
pub fn new(name: &str, value: &str) -> ZaiResult<Self> {
let Some(static_name) = ALLOWED_HEADER_NAMES
.iter()
.copied()
.find(|candidate| candidate.eq_ignore_ascii_case(name))
else {
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.bytes().all(|byte| (0x20..=0x7e).contains(&byte)) {
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: "additional header value must contain printable ASCII only".to_string(),
});
}
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 validate(&self) -> ZaiResult<()> {
if self.connect_timeout.is_zero() || self.connect_timeout > Duration::from_secs(10) {
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: "connect_timeout must be in 1ns..=10s".to_string(),
});
}
if self.request_timeout.is_zero() || self.request_timeout > Duration::from_secs(60) {
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: "request_timeout must be in 1ns..=60s".to_string(),
});
}
if !(1..=3).contains(&self.max_attempts) {
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: "max_attempts must be 1, 2 or 3".to_string(),
});
}
let mut names = std::collections::HashSet::with_capacity(self.additional_headers.len());
for header in &self.additional_headers {
if !names.insert(header.name()) {
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: format!(
"additional header {:?} must not be configured more than once",
header.name()
),
});
}
}
Ok(())
}
pub fn with_request_timeout(mut self, d: Duration) -> ZaiResult<Self> {
if d.is_zero() || d > Duration::from_secs(60) {
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: "request_timeout must be in 1ns..=60s".to_string(),
});
}
self.request_timeout = d;
Ok(self)
}
pub fn with_connect_timeout(mut self, d: Duration) -> ZaiResult<Self> {
if d.is_zero() || d > Duration::from_secs(10) {
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_CONFIG,
message: "connect_timeout must be in 1ns..=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
}
}
#[derive(Debug, Clone)]
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 builder_rejects_keys_that_cannot_be_safe_header_credentials() {
assert!(ZaiClient::builder("abc.def\nghi").build().is_err());
assert!(ZaiClient::builder("abc.def ghi").build().is_err());
assert!(ZaiClient::builder("密钥.abcdefghij").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_eq!(
AdditionalHeader::new("x-test-client", "preserved")
.unwrap()
.name(),
"X-Test-Client"
);
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());
assert!(AdditionalHeader::new("X-Test-Client", "非 ASCII").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());
assert!(
HttpTransportConfig::default()
.with_request_timeout(Duration::ZERO)
.is_err()
);
}
#[test]
fn client_build_validates_direct_transport_fields() {
let invalid = HttpTransportConfig {
max_attempts: 0,
..HttpTransportConfig::default()
};
assert!(
ZaiClient::builder("abcdefghij.0123456789abcdef")
.transport(invalid)
.build()
.is_err()
);
assert!(
ZaiClient::builder(" abcdefghij.0123456789abcdef ")
.build()
.is_err()
);
let duplicate_headers = HttpTransportConfig::default()
.with_additional_header(AdditionalHeader::new("X-Test-Client", "a").unwrap())
.with_additional_header(AdditionalHeader::new("x-test-client", "b").unwrap());
assert!(duplicate_headers.validate().is_err());
}
}