use serde::{Deserialize, Deserializer, Serialize, Serializer};
use validator::Validate;
use super::types::KnowledgeResponse;
use crate::ZaiResult;
use crate::client::ZaiClient;
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum EmbeddingId {
Embedding2,
Embedding3New,
Embedding3Pro,
}
impl EmbeddingId {
pub fn as_i64(&self) -> i64 {
match self {
EmbeddingId::Embedding2 => 3,
EmbeddingId::Embedding3New => 11,
EmbeddingId::Embedding3Pro => 12,
}
}
pub fn as_model_name(&self) -> &'static str {
match self {
EmbeddingId::Embedding2 => "Embedding-2",
EmbeddingId::Embedding3New => "Embedding-3",
EmbeddingId::Embedding3Pro => "Embedding-3-pro",
}
}
}
impl Serialize for EmbeddingId {
fn serialize<S: Serializer>(&self, serializer: S) -> Result<S::Ok, S::Error> {
serializer.serialize_i64(self.as_i64())
}
}
impl<'de> Deserialize<'de> for EmbeddingId {
fn deserialize<D: Deserializer<'de>>(deserializer: D) -> Result<Self, D::Error> {
let v = i64::deserialize(deserializer)?;
match v {
3 => Ok(EmbeddingId::Embedding2),
11 => Ok(EmbeddingId::Embedding3New),
12 => Ok(EmbeddingId::Embedding3Pro),
other => Err(serde::de::Error::custom(format!(
"unsupported embedding_id: {other} (expected 3, 11 or 12)"
))),
}
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum BackgroundColor {
Blue,
Red,
Orange,
Purple,
Sky,
Green,
Yellow,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum KnowledgeIcon {
Question,
Book,
Seal,
Wrench,
Tag,
Horn,
House,
}
#[derive(Clone, Serialize, Deserialize, Validate)]
pub struct KnowledgeCreateBody {
pub embedding_id: EmbeddingId,
#[validate(length(min = 1))]
pub name: String,
#[serde(skip_serializing_if = "Option::is_none")]
pub description: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub background: Option<BackgroundColor>,
#[serde(skip_serializing_if = "Option::is_none")]
pub icon: Option<KnowledgeIcon>,
#[serde(skip_serializing_if = "Option::is_none")]
pub embedding_model: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
#[validate(range(max = 1))]
pub contextual: Option<u8>,
}
impl std::fmt::Debug for KnowledgeCreateBody {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter
.debug_struct("KnowledgeCreateBody")
.field("embedding_id", &self.embedding_id)
.field("name", &"[REDACTED]")
.field("description_configured", &self.description.is_some())
.field("background", &self.background)
.field("icon", &self.icon)
.field(
"embedding_model_configured",
&self.embedding_model.is_some(),
)
.field("contextual", &self.contextual)
.finish()
}
}
impl KnowledgeCreateBody {
fn validate_embedding_pair(&self) -> ZaiResult<()> {
if let Some(model) = self.embedding_model.as_deref()
&& model != self.embedding_id.as_model_name()
{
return Err(crate::ZaiError::ApiError {
code: crate::client::error::codes::SDK_VALIDATION,
message: format!(
"embedding_id {} requires embedding_model '{}'",
self.embedding_id.as_i64(),
self.embedding_id.as_model_name()
),
});
}
Ok(())
}
}
pub struct KnowledgeCreateRequest {
body: KnowledgeCreateBody,
}
impl KnowledgeCreateRequest {
pub fn new(embedding_id: EmbeddingId, name: impl Into<String>) -> Self {
let body = KnowledgeCreateBody {
embedding_id,
name: name.into(),
description: None,
background: None,
icon: None,
embedding_model: None,
contextual: None,
};
Self { body }
}
pub fn with_description(mut self, desc: impl Into<String>) -> Self {
self.body.description = Some(desc.into());
self
}
pub fn with_background(mut self, bg: BackgroundColor) -> Self {
self.body.background = Some(bg);
self
}
pub fn with_icon(mut self, icon: KnowledgeIcon) -> Self {
self.body.icon = Some(icon);
self
}
pub fn with_embedding_model(mut self, model: impl Into<String>) -> Self {
self.body.embedding_model = Some(model.into());
self
}
pub fn with_contextual(mut self, contextual: u8) -> Self {
self.body.contextual = Some(contextual);
self
}
pub fn validate(&self) -> ZaiResult<()> {
self.body.validate()?;
if self.body.name.trim().is_empty() {
return Err(crate::client::validation::invalid(
"knowledge name must not be blank",
));
}
self.body.validate_embedding_pair()
}
pub async fn send_via(&self, client: &ZaiClient) -> ZaiResult<KnowledgeCreateResponse> {
self.validate()?;
let route = crate::client::routes::KNOWLEDGE_CREATE;
let url = client.endpoints().resolve_route(route, &[])?;
client
.send_json::<_, KnowledgeCreateResponse>(route.method(), url, &self.body)
.await
}
}
#[derive(Debug, Clone, Serialize, Deserialize, Validate)]
pub struct KnowledgeCreateData {
#[serde(skip_serializing_if = "Option::is_none")]
pub id: Option<String>,
}
pub type KnowledgeCreateResponse = KnowledgeResponse<KnowledgeCreateData>;
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn rejects_blank_names_and_mismatched_embedding_models() {
assert!(
KnowledgeCreateRequest::new(EmbeddingId::Embedding3New, " \t")
.validate()
.is_err()
);
assert!(
KnowledgeCreateRequest::new(EmbeddingId::Embedding3New, "docs")
.with_embedding_model("Embedding-2")
.validate()
.is_err()
);
assert!(
KnowledgeCreateRequest::new(EmbeddingId::Embedding3New, "docs")
.with_embedding_model("Embedding-3")
.validate()
.is_ok()
);
}
#[test]
fn request_body_debug_redacts_names_and_descriptions() {
let body = KnowledgeCreateBody {
embedding_id: EmbeddingId::Embedding3New,
name: "private-name".to_owned(),
description: Some("private-description".to_owned()),
background: Some(BackgroundColor::Blue),
icon: Some(KnowledgeIcon::Book),
embedding_model: Some("private-model-string".to_owned()),
contextual: Some(1),
};
let debug = format!("{body:?}");
for secret in [
"private-name",
"private-description",
"private-model-string",
] {
assert!(!debug.contains(secret));
}
}
}