use derive_builder::Builder;
use dynamo_tokens::{Token, TokenBlockMmInfo, validate_and_sort_mm_info};
use serde::{Deserialize, Serialize};
use crate::error::KvHashingError;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct RequestMmObjectInfo {
pub mm_hash: u64,
pub offset: usize,
pub length: usize,
}
impl From<RequestMmObjectInfo> for TokenBlockMmInfo {
fn from(v: RequestMmObjectInfo) -> Self {
Self {
mm_hash: v.mm_hash,
offset: v.offset,
length: v.length,
}
}
}
impl From<TokenBlockMmInfo> for RequestMmObjectInfo {
fn from(v: TokenBlockMmInfo) -> Self {
Self {
mm_hash: v.mm_hash,
offset: v.offset,
length: v.length,
}
}
}
#[derive(Debug, Clone, Builder)]
#[builder(
pattern = "owned",
build_fn(private, name = "build_internal", error = "KvHashingError"),
derive(Debug)
)]
pub struct Request {
#[builder(setter(into))]
pub(crate) tokens: Vec<Token>,
#[builder(default, setter(into))]
pub(crate) lora_name: Option<String>,
#[builder(default, setter(into))]
pub(crate) salt: Option<String>,
#[builder(default)]
pub(crate) mm_info: Vec<RequestMmObjectInfo>,
}
impl Request {
pub fn builder() -> RequestBuilder {
RequestBuilder::default()
}
pub fn tokens(&self) -> &[Token] {
&self.tokens
}
pub fn lora_name(&self) -> Option<&str> {
self.lora_name.as_deref()
}
pub fn salt(&self) -> Option<&str> {
self.salt.as_deref()
}
pub fn mm_info(&self) -> &[RequestMmObjectInfo] {
&self.mm_info
}
pub(crate) fn token_mm_info(&self) -> Vec<TokenBlockMmInfo> {
self.mm_info.iter().copied().map(Into::into).collect()
}
}
impl RequestBuilder {
pub fn build(self) -> Result<Request, KvHashingError> {
let mut request = self.build_internal()?;
let token_mm: Vec<TokenBlockMmInfo> = std::mem::take(&mut request.mm_info)
.into_iter()
.map(Into::into)
.collect();
let validated = validate_and_sort_mm_info(&token_mm, request.tokens.len())?;
request.mm_info = validated.into_iter().map(Into::into).collect();
Ok(request)
}
}
impl From<derive_builder::UninitializedFieldError> for KvHashingError {
fn from(e: derive_builder::UninitializedFieldError) -> Self {
KvHashingError::MissingField(e.field_name())
}
}