use std::collections::BTreeMap;
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use crate::records::byte_format;
use crate::records::identifiers::{
Capability, ExecutionMode, Modality, ModelState, RunTier, RuntimeId, SourceKind,
};
use crate::records::json_value::JsonValue;
use crate::time::now_millis;
#[derive(Debug, Clone, PartialEq, Eq, Hash, Serialize, Deserialize)]
pub struct ModelSource {
pub kind: SourceKind,
pub path: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub repo: Option<String>,
#[serde(rename = "ref", default, skip_serializing_if = "Option::is_none")]
pub reference: Option<String>,
}
impl ModelSource {
pub fn new(kind: SourceKind, path: &str) -> Self {
Self {
kind,
path: path.to_owned(),
repo: None,
reference: None,
}
}
pub fn identity(&self) -> String {
format!(
"{}|{}|{}",
self.kind.as_str(),
self.path,
self.repo.as_deref().unwrap_or("")
)
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum Resolution {
Auto,
User,
#[default]
Unresolved,
}
#[derive(Debug, Clone, PartialEq, Eq, Default, Serialize, Deserialize)]
pub struct RuntimeRef {
#[serde(default, skip_serializing_if = "Option::is_none")]
pub id: Option<RuntimeId>,
#[serde(default)]
pub resolved: Resolution,
#[serde(default)]
pub tier: RunTier,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub alternatives: Vec<RuntimeId>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub confirmed_at: Option<i64>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)]
#[serde(rename_all = "lowercase")]
pub enum ParamType {
Int,
Float,
Bool,
String,
Enum,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ParamSpec {
pub key: String,
#[serde(rename = "type")]
pub param_type: ParamType,
#[serde(rename = "default", default, skip_serializing_if = "Option::is_none")]
pub default_value: Option<JsonValue>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub range: Option<Vec<JsonValue>>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub values: Option<Vec<String>>,
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ModelRecord {
pub id: String,
pub name: String,
pub modality: Modality,
pub capabilities: Vec<Capability>,
pub source: ModelSource,
#[serde(default)]
pub runtime: RuntimeRef,
#[serde(default)]
pub params: Vec<ParamSpec>,
#[serde(default, skip_serializing_if = "BTreeMap::is_empty")]
pub param_values: BTreeMap<String, JsonValue>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub system_prompt: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub alias: Option<String>,
#[serde(default)]
pub execution: ExecutionMode,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub footprint_bytes: Option<i64>,
#[serde(default, rename = "footprint_mb", skip_serializing)]
pub(crate) legacy_footprint_mb: Option<i64>,
#[serde(default)]
pub state: ModelState,
pub registered_at: i64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub primary_weight_path: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub context_length: Option<i64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub has_chat_template: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub supports_tools: Option<bool>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub stop_tokens: Option<Vec<String>>,
#[serde(default, skip_serializing_if = "std::ops::Not::not")]
pub downloading: bool,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub content_fingerprint: Option<String>,
}
impl ModelRecord {
pub fn new(
name: &str,
modality: Modality,
capabilities: Vec<Capability>,
source: ModelSource,
) -> Self {
Self {
id: stable_id(&source),
name: name.to_owned(),
modality,
capabilities,
source,
runtime: RuntimeRef::default(),
params: Vec::new(),
param_values: BTreeMap::new(),
system_prompt: None,
alias: None,
execution: ExecutionMode::Sync,
footprint_bytes: None,
legacy_footprint_mb: None,
state: ModelState::Unresolved,
registered_at: now_millis(),
primary_weight_path: None,
context_length: None,
has_chat_template: None,
supports_tools: None,
stop_tokens: None,
downloading: false,
content_fingerprint: None,
}
}
pub fn display_name(&self) -> &str {
match &self.alias {
Some(alias) if !alias.is_empty() => alias,
_ => &self.name,
}
}
pub fn wire_id(&self) -> &str {
self.display_name()
}
pub fn size_on_disk(&self) -> Option<i64> {
self.footprint_bytes.filter(|bytes| *bytes > 0)
}
pub(crate) fn adopt_legacy_footprint(&mut self) {
if let Some(mebibytes) = self.legacy_footprint_mb.take()
&& self.footprint_bytes.is_none()
{
self.footprint_bytes = Some(mebibytes.saturating_mul(byte_format::BYTES_PER_MIB));
}
}
pub fn footprint_mib(&self) -> Option<i64> {
self.size_on_disk()
.map(|bytes| bytes / byte_format::BYTES_PER_MIB)
}
pub fn can(&self, capability: &Capability) -> bool {
self.capabilities.contains(capability)
}
}
pub fn stable_id(source: &ModelSource) -> String {
let mut hasher = Sha256::new();
for field in [
source.kind.as_str(),
source.path.as_str(),
source.repo.as_deref().unwrap_or(""),
] {
hasher.update((field.len() as u64).to_le_bytes());
hasher.update(field.as_bytes());
}
let digest = hasher.finalize();
hex::encode(&digest[..8])
}