diff --git a/src/app.rs b/src/app.rs index ead7d6a..58ee90a 100644 --- a/src/app.rs +++ b/src/app.rs @@ -1,7 +1,17 @@ +mod generation; +mod model_manager; +mod preferences; +mod projects; mod view; pub(crate) use view::app_theme; +use generation::ChatMessage; +use model_manager::{ActiveDownload, ModelDownload, ModelOperation}; +use preferences::PreferenceDraft; +#[cfg(test)] +use preferences::{parse_optional_gib, parse_streaming_cache}; + use crate::database::{AppPreferences, Database, ProjectWithSessions, StoredMessage}; #[cfg(target_os = "macos")] use crate::engine::ChatTurn; @@ -29,368 +39,6 @@ use std::time::{Duration, Instant}; const APP_ID: &str = "DS4Server.rfc1437.de"; const METRICS_SAMPLE_INTERVAL: Duration = Duration::from_millis(200); -#[derive(Clone)] -struct PreferenceDraft { - model: ModelChoice, - dspark_enabled: bool, - idle_timeout_minutes: String, - endpoint_port: String, - context_tokens: String, - max_generated_tokens: String, - system_prompt: String, - temperature: String, - top_p: String, - min_p: String, - seed: String, - reasoning_mode: ReasoningMode, - cpu_threads: String, - power_percent: String, - prefill_chunk: String, - quality: bool, - warm_weights: bool, - mtp_draft_tokens: String, - mtp_margin: String, - glm_mtp: bool, - glm_mtp_timing: bool, - dspark_confidence_threshold: String, - dspark_strict: bool, - ssd_streaming: bool, - ssd_streaming_cold: bool, - ssd_cache: String, - ssd_full_layers: String, - ssd_preload_experts: String, - directional_steering_file: String, - directional_steering_ffn: String, - directional_steering_attn: String, - simulated_used_memory_gib: String, - expert_profile_path: String, -} - -impl PreferenceDraft { - fn from_saved(preferences: &AppPreferences) -> Result { - let model = ModelChoice::from_id(&preferences.selected_model) - .ok_or_else(|| format!("Unsupported model: {}", preferences.selected_model))?; - let generation = preferences.generation()?; - let runtime = preferences.runtime()?; - runtime.validate(model)?; - let execution = &runtime.execution; - let speculative = &runtime.speculative; - Ok(Self { - model, - dspark_enabled: speculative.dspark_enabled, - idle_timeout_minutes: preferences.idle_timeout_minutes.to_string(), - endpoint_port: preferences.endpoint_port.to_string(), - context_tokens: generation.context_tokens.to_string(), - max_generated_tokens: generation.max_generated_tokens.to_string(), - system_prompt: generation.system_prompt, - temperature: generation - .temperature - .map_or_else(String::new, |value| value.to_string()), - top_p: generation - .top_p - .map_or_else(String::new, |value| value.to_string()), - min_p: generation - .min_p - .map_or_else(String::new, |value| value.to_string()), - seed: generation - .seed - .map_or_else(String::new, |value| value.to_string()), - reasoning_mode: generation.reasoning_mode, - cpu_threads: execution - .cpu_threads - .map_or_else(String::new, |value| value.to_string()), - power_percent: execution - .power_percent - .map_or_else(String::new, |value| value.to_string()), - prefill_chunk: execution - .prefill_chunk - .map_or_else(String::new, |value| value.to_string()), - quality: execution.quality, - warm_weights: execution.warm_weights, - mtp_draft_tokens: speculative.mtp_draft_tokens.to_string(), - mtp_margin: speculative.mtp_margin.to_string(), - glm_mtp: speculative.glm_mtp, - glm_mtp_timing: speculative.glm_mtp_timing, - dspark_confidence_threshold: speculative - .dspark_confidence_threshold - .map_or_else(String::new, |value| value.to_string()), - dspark_strict: speculative.dspark_strict, - ssd_streaming: runtime.ssd.enabled, - ssd_streaming_cold: runtime.ssd.cold, - ssd_cache: runtime - .ssd - .cache - .map_or_else(String::new, |cache| match cache { - StreamingCacheBudget::Experts(experts) => experts.to_string(), - StreamingCacheBudget::Gib(gib) => format!("{gib}GB"), - }), - ssd_full_layers: runtime - .ssd - .full_layers - .map_or_else(String::new, |value| value.to_string()), - ssd_preload_experts: runtime - .ssd - .preload_experts - .map_or_else(String::new, |value| value.to_string()), - directional_steering_file: runtime.steering.file.unwrap_or_default(), - directional_steering_ffn: runtime - .steering - .ffn_scale - .map_or_else(String::new, |value| value.to_string()), - directional_steering_attn: runtime - .steering - .attention_scale - .map_or_else(String::new, |value| value.to_string()), - simulated_used_memory_gib: runtime - .diagnostics - .simulated_used_memory_gib - .map_or_else(String::new, |value| value.to_string()), - expert_profile_path: runtime.diagnostics.expert_profile_path.unwrap_or_default(), - }) - } - - fn generation(&self) -> Result { - let preferences = GenerationPreferences { - context_tokens: parse_positive_i32("Context tokens", &self.context_tokens)?, - max_generated_tokens: parse_positive_i32( - "Maximum generated tokens", - &self.max_generated_tokens, - )?, - system_prompt: self.system_prompt.clone(), - temperature: parse_optional_f32("Temperature", &self.temperature)?, - top_p: parse_optional_f32("Top-p", &self.top_p)?, - min_p: parse_optional_f32("Min-p", &self.min_p)?, - seed: parse_optional_u64("Seed", &self.seed)?, - reasoning_mode: self.reasoning_mode, - }; - preferences.validate()?; - Ok(preferences) - } - - fn reset(&mut self) { - let defaults = GenerationPreferences::default(); - let execution = ExecutionPreferences::default(); - let speculative = SpeculativePreferences::default(); - let ssd = SsdPreferences::default(); - self.model = ModelChoice::default(); - self.dspark_enabled = false; - self.idle_timeout_minutes = "10".into(); - self.endpoint_port = "4000".into(); - self.context_tokens = defaults.context_tokens.to_string(); - self.max_generated_tokens = defaults.max_generated_tokens.to_string(); - self.system_prompt = defaults.system_prompt; - self.temperature.clear(); - self.top_p.clear(); - self.min_p.clear(); - self.seed.clear(); - self.reasoning_mode = defaults.reasoning_mode; - self.cpu_threads.clear(); - self.power_percent.clear(); - self.prefill_chunk.clear(); - self.quality = execution.quality; - self.warm_weights = execution.warm_weights; - self.mtp_draft_tokens = speculative.mtp_draft_tokens.to_string(); - self.mtp_margin = speculative.mtp_margin.to_string(); - self.glm_mtp = speculative.glm_mtp; - self.glm_mtp_timing = speculative.glm_mtp_timing; - self.dspark_confidence_threshold.clear(); - self.dspark_strict = speculative.dspark_strict; - self.ssd_streaming = ssd.enabled; - self.ssd_streaming_cold = ssd.cold; - self.ssd_cache.clear(); - self.ssd_full_layers.clear(); - self.ssd_preload_experts.clear(); - self.directional_steering_file.clear(); - self.directional_steering_ffn.clear(); - self.directional_steering_attn.clear(); - self.simulated_used_memory_gib.clear(); - self.expert_profile_path.clear(); - } - - fn execution(&self) -> Result { - Ok(ExecutionPreferences { - cpu_threads: parse_optional_u32("CPU helper threads", &self.cpu_threads)?, - power_percent: parse_optional_u8("GPU power", &self.power_percent)?, - prefill_chunk: parse_optional_u32("Prefill chunk", &self.prefill_chunk)?, - quality: self.quality, - warm_weights: self.warm_weights, - }) - } - - fn speculative(&self) -> Result { - Ok(SpeculativePreferences { - mtp_draft_tokens: parse_positive_i32("MTP draft tokens", &self.mtp_draft_tokens)?, - mtp_margin: parse_f32("MTP margin", &self.mtp_margin)?, - glm_mtp: self.glm_mtp, - glm_mtp_timing: self.glm_mtp_timing, - dspark_enabled: self.dspark_enabled, - dspark_confidence_threshold: parse_optional_f32( - "DSpark confidence", - &self.dspark_confidence_threshold, - )?, - dspark_strict: self.dspark_strict, - }) - } - - fn runtime(&self) -> Result { - Ok(RuntimePreferences { - execution: self.execution()?, - speculative: self.speculative()?, - ssd: SsdPreferences { - enabled: self.ssd_streaming, - cold: self.ssd_streaming_cold, - cache: parse_streaming_cache(&self.ssd_cache)?, - full_layers: parse_optional_u32("SSD full-layer count", &self.ssd_full_layers)?, - preload_experts: parse_optional_u32( - "SSD preload experts", - &self.ssd_preload_experts, - )?, - }, - steering: SteeringPreferences { - file: optional_text(&self.directional_steering_file), - ffn_scale: parse_optional_f32( - "Directional FFN scale", - &self.directional_steering_ffn, - )?, - attention_scale: parse_optional_f32( - "Directional attention scale", - &self.directional_steering_attn, - )?, - }, - diagnostics: DiagnosticPreferences { - simulated_used_memory_gib: parse_optional_gib( - "Simulated used memory", - &self.simulated_used_memory_gib, - )?, - expert_profile_path: optional_text(&self.expert_profile_path), - }, - }) - } -} - -fn parse_positive_i32(name: &str, value: &str) -> Result { - value - .trim() - .parse::() - .ok() - .filter(|value| *value > 0) - .ok_or_else(|| format!("{name} must be a positive whole number.")) -} - -fn parse_optional_f32(name: &str, value: &str) -> Result, String> { - let value = value.trim(); - if value.is_empty() { - Ok(None) - } else { - value - .parse() - .map(Some) - .map_err(|_| format!("{name} must be a number or left blank for the DS4 default.")) - } -} - -fn parse_f32(name: &str, value: &str) -> Result { - value - .trim() - .parse() - .map_err(|_| format!("{name} must be a number.")) -} - -fn parse_optional_u64(name: &str, value: &str) -> Result, String> { - let value = value.trim(); - if value.is_empty() { - Ok(None) - } else { - value - .parse() - .map(Some) - .map_err(|_| format!("{name} must be a positive whole number or left blank.")) - } -} - -fn parse_optional_u32(name: &str, value: &str) -> Result, String> { - parse_optional_number(name, value) -} - -fn parse_optional_u8(name: &str, value: &str) -> Result, String> { - parse_optional_number(name, value) -} - -fn parse_optional_number( - name: &str, - value: &str, -) -> Result, String> { - let value = value.trim(); - if value.is_empty() { - Ok(None) - } else { - value - .parse() - .map(Some) - .map_err(|_| format!("{name} must be a positive whole number or left blank.")) - } -} - -fn parse_streaming_cache(value: &str) -> Result, String> { - let value = value.trim(); - if value.is_empty() { - return Ok(None); - } - if value.len() > 2 - && value - .get(value.len() - 2..) - .is_some_and(|suffix| suffix.eq_ignore_ascii_case("gb")) - { - return parse_gib("SSD cache budget", value) - .map(|gib| Some(StreamingCacheBudget::Gib(gib))); - } - if !value.chars().all(|character| character.is_ascii_digit()) { - return Err( - "SSD cache budget must be a positive expert count or whole GiB value such as 64GB." - .into(), - ); - } - value - .parse::() - .ok() - .filter(|value| *value > 0) - .map(StreamingCacheBudget::Experts) - .map(Some) - .ok_or_else(|| { - "SSD cache budget must be a positive expert count or whole GiB value such as 64GB." - .into() - }) -} - -fn parse_optional_gib(name: &str, value: &str) -> Result, String> { - let value = value.trim(); - if value.is_empty() { - Ok(None) - } else { - parse_gib(name, value).map(Some) - } -} - -fn parse_gib(name: &str, value: &str) -> Result { - let value = value - .get(value.len().saturating_sub(2)..) - .filter(|suffix| suffix.eq_ignore_ascii_case("gb")) - .map_or(value, |_| &value[..value.len() - 2]); - if !value.chars().all(|character| character.is_ascii_digit()) { - return Err(format!("{name} must be a positive whole GiB value.")); - } - value - .parse::() - .ok() - .filter(|value| *value > 0 && *value <= u64::MAX / GIB) - .ok_or_else(|| format!("{name} must be a positive whole GiB value.")) -} - -fn optional_text(value: &str) -> Option { - let value = value.trim(); - (!value.is_empty()).then(|| value.to_owned()) -} - pub(crate) struct App { main_window: window::Id, pub(super) model_manager_window: Option, @@ -449,83 +97,6 @@ pub(super) struct MetricsPoint { pub(super) kv_write_bytes_per_second: f32, } -#[derive(Clone, Debug)] -pub(super) struct ChatMessage { - id: i32, - pub(super) user: bool, - pub(super) reasoning: Option, - reasoning_complete: bool, - pub(super) reasoning_open: bool, - pub(super) content: String, - pub(super) markdown: Vec, -} - -impl ChatMessage { - fn append(&mut self, reasoning: bool, chunk: &str) { - if reasoning { - self.reasoning.get_or_insert_default().push_str(chunk); - } else { - self.reasoning_complete = true; - self.content.push_str(chunk); - } - } - - fn refresh_markdown(&mut self) { - if !self.user { - let content = if self.reasoning.is_some() { - self.content.trim_start() - } else { - &self.content - }; - self.markdown = markdown::parse(content).collect(); - } - } -} - -impl From for ChatMessage { - fn from(message: StoredMessage) -> Self { - let mut message = Self { - id: message.id, - user: message.user, - reasoning: message.reasoning, - reasoning_complete: message.reasoning_complete, - reasoning_open: false, - content: message.content, - markdown: Vec::new(), - }; - message.refresh_markdown(); - message - } -} - -#[derive(Debug)] -pub(super) enum ModelDownload { - Idle, - Active(ActiveDownload), - Complete(ManagedArtifactId, ModelOperation, DownloadProgress), - Failed(ManagedArtifactId, String, DownloadProgress), -} - -#[derive(Clone, Copy, Debug, Eq, PartialEq)] -pub(super) enum ModelOperation { - Download, - Validate, -} - -#[derive(Debug)] -pub(super) struct ActiveDownload { - artifact: ManagedArtifactId, - operation: ModelOperation, - progress: DownloadProgress, - sampled_at: Instant, - sampled_bytes: u64, - bytes_per_second: f64, - verified_bytes: Arc, - cancel: Arc, - result: mpsc::Receiver>, - stopping: bool, -} - #[derive(Debug, Clone)] pub(crate) enum Message { Noop, @@ -1234,552 +805,6 @@ impl App { } } - fn open_model_manager(&mut self) -> Task { - if let Some(id) = self.model_manager_window { - return window::gain_focus(id); - } - let (id, open) = window::open(window::Settings { - size: Size::new(760.0, 560.0), - min_size: Some(Size::new(620.0, 420.0)), - icon: Some(app_icon()), - ..Default::default() - }); - self.model_manager_window = Some(id); - open.map(Message::ModelManagerOpened) - } - - fn start_model_operation(&mut self, artifact: ManagedArtifactId, operation: ModelOperation) { - if matches!(self.model_download, ModelDownload::Active(_)) { - return; - } - let models_path = models_path(); - let progress = match operation { - ModelOperation::Download => model::artifact_download_progress(artifact, &models_path), - ModelOperation::Validate => model::artifact_verification_progress(artifact, 0), - }; - let cancel = Arc::new(AtomicBool::new(false)); - let worker_cancel = Arc::clone(&cancel); - let verified_bytes = Arc::new(AtomicU64::new(0)); - let worker_verified_bytes = Arc::clone(&verified_bytes); - let (result_sender, result_receiver) = mpsc::channel(); - let thread_name = match operation { - ModelOperation::Download => "model-download", - ModelOperation::Validate => "model-validation", - }; - if let Err(error) = thread::Builder::new() - .name(thread_name.to_owned()) - .spawn(move || { - let result = match operation { - ModelOperation::Download => model::download_managed_artifact( - artifact, - &models_path, - &worker_cancel, - &worker_verified_bytes, - ), - ModelOperation::Validate => model::validate_managed_artifact( - artifact, - &models_path, - &worker_cancel, - &worker_verified_bytes, - ), - }; - let _ = result_sender.send(result); - }) - { - self.error = Some(format!("Could not start {thread_name}: {error}")); - return; - } - self.pending_model_delete = None; - self.model_download = ModelDownload::Active(ActiveDownload { - artifact, - operation, - sampled_at: Instant::now(), - sampled_bytes: progress.completed(), - progress, - bytes_per_second: 0.0, - verified_bytes, - cancel, - result: result_receiver, - stopping: false, - }); - } - - fn open_preferences(&mut self) { - if self.database.is_none() || self.pending_project_path.is_some() || self.choosing_folder { - return; - } - match PreferenceDraft::from_saved(&self.preferences) { - Ok(draft) => { - self.preference_draft = draft; - self.preference_error = None; - self.preferences_open = true; - } - Err(error) => self.error = Some(error), - } - } - - fn save_preferences(&mut self) { - let Ok(idle_timeout_minutes) = self - .preference_draft - .idle_timeout_minutes - .trim() - .parse::() - else { - self.preference_error = Some("Idle timeout must be a whole number.".into()); - return; - }; - if !(1..=1440).contains(&idle_timeout_minutes) { - self.preference_error = Some("Idle timeout must be between 1 and 1440 minutes.".into()); - return; - } - let Ok(endpoint_port) = self.preference_draft.endpoint_port.trim().parse::() else { - self.preference_error = Some("Endpoint port must be a whole number.".into()); - return; - }; - if endpoint_port == 0 { - self.preference_error = Some("Endpoint port must be between 1 and 65535.".into()); - return; - } - let generation = match self.preference_draft.generation() { - Ok(generation) => generation, - Err(error) => { - self.preference_error = Some(error); - return; - } - }; - - let model = self.preference_draft.model; - let runtime = match self.preference_draft.runtime() { - Ok(runtime) => runtime, - Err(error) => { - self.preference_error = Some(error); - return; - } - }; - if let Err(error) = runtime.validate(model) { - self.preference_error = Some(error); - return; - } - #[cfg(target_os = "macos")] - let pending_endpoint = if self.preferences.endpoint_port != i32::from(endpoint_port) - || self._endpoint.is_none() - { - let Some(generation) = &self.generation_service else { - self.preference_error = Some("The model runtime is unavailable.".into()); - return; - }; - match crate::server::ServerHandle::spawn( - generation.clone(), - Arc::clone(&self.runtime_preferences), - models_path(), - application_support_path().join("kv-cache").join("http"), - endpoint_port, - Arc::clone(&self.metrics), - ) { - Ok(endpoint) => Some(endpoint), - Err(error) => { - self.preference_error = Some(error); - return; - } - } - } else { - None - }; - let Some(database) = &mut self.database else { - return; - }; - match database.update_preferences( - model.id(), - idle_timeout_minutes, - i32::from(endpoint_port), - &generation, - &runtime, - ) { - Ok(preferences) => { - self.preferences = preferences; - #[cfg(target_os = "macos")] - if let Ok(mut runtime_preferences) = self.runtime_preferences.write() { - *runtime_preferences = self.preferences.clone(); - } - #[cfg(target_os = "macos")] - if let Some(endpoint) = pending_endpoint { - self._endpoint = Some(endpoint); - } - self.preference_draft = PreferenceDraft::from_saved(&self.preferences) - .expect("the saved model was selected from the supported catalog"); - self.preferences_open = false; - self.preference_error = None; - self.error = None; - } - Err(error) => self.preference_error = Some(error), - } - } - - fn prepare_project(&mut self, path: PathBuf) { - let Ok(path) = fs::canonicalize(path) else { - self.error = Some("The selected folder is no longer available.".into()); - return; - }; - let Some(path_text) = path.to_str() else { - self.error = Some("The selected folder path is not valid UTF-8.".into()); - return; - }; - if self - .projects - .iter() - .any(|item| item.project.path == path_text) - { - self.error = Some("That project is already in the sidebar.".into()); - return; - } - - self.project_name_input = path - .file_name() - .and_then(|name| name.to_str()) - .unwrap_or("Project") - .to_owned(); - self.pending_project_path = Some(path); - self.error = None; - } - - fn create_project(&mut self) { - let name = self.project_name_input.trim(); - if name.is_empty() { - self.error = Some("Project name cannot be empty.".into()); - return; - } - let Some(path) = &self.pending_project_path else { - return; - }; - let Some(path) = path.to_str() else { - self.error = Some("The selected folder path is not valid UTF-8.".into()); - return; - }; - let Some(database) = &mut self.database else { - return; - }; - - match database.create_project(name, path) { - Ok(project) => { - self.selected_project = Some(project.id); - self.selected_session = None; - self.pending_project_path = None; - self.project_name_input.clear(); - self.error = None; - self.reload_projects(); - } - Err(error) => self.error = Some(error), - } - } - - fn create_session(&mut self) { - if self.generating { - self.error = Some("Stop the active generation before creating a session.".into()); - return; - } - let Some(project_id) = self.selected_project else { - self.error = Some("Select a project first.".into()); - return; - }; - let default_number = self - .selected_project() - .map(|project| project.sessions.len() + 1) - .unwrap_or(1); - let title = format!("Session {default_number}"); - let Some(database) = &mut self.database else { - return; - }; - - match database.create_session(project_id, &title) { - Ok(session) => { - self.selected_session = Some(session.id); - self.conversation.clear(); - self.composer.clear(); - self.context_used = 0; - self.context_limit = self.preferences.context_tokens.max(0) as u32; - self.tokens_per_second = None; - self.error = None; - self.reload_projects(); - } - Err(error) => self.error = Some(error), - } - } - - fn reload_projects(&mut self) { - if let Some(database) = &mut self.database { - match database.load_projects() { - Ok(projects) => self.projects = projects, - Err(error) => self.error = Some(error), - } - } - } - - fn update_download_progress(&mut self) { - let ModelDownload::Active(download) = &mut self.model_download else { - return; - }; - let now = Instant::now(); - let verified = download.verified_bytes.load(Ordering::Relaxed); - let mut progress = match download.operation { - ModelOperation::Download => { - model::artifact_download_progress(download.artifact, &models_path()) - } - ModelOperation::Validate => { - model::artifact_verification_progress(download.artifact, verified) - } - }; - if download.operation == ModelOperation::Download - && let Some(verification) = &mut progress.verification - { - verification.verified = verified.min(verification.total); - } - let elapsed = now.duration_since(download.sampled_at).as_secs_f64(); - let completed = progress.completed(); - let phase_changed = progress.phase != download.progress.phase; - let transferred = completed.saturating_sub(download.sampled_bytes); - if phase_changed { - download.bytes_per_second = 0.0; - } else if transferred > 0 && elapsed > 0.0 { - let current = transferred as f64 / elapsed; - download.bytes_per_second = if download.bytes_per_second == 0.0 { - current - } else { - download.bytes_per_second * 0.75 + current * 0.25 - }; - } - download.progress = progress; - download.sampled_at = now; - download.sampled_bytes = completed; - - let result = match download.result.try_recv() { - Ok(result) => Some(result), - Err(TryRecvError::Empty) => None, - Err(TryRecvError::Disconnected) => { - Some(Err("Download worker stopped unexpectedly.".into())) - } - }; - let artifact = download.artifact; - let operation = download.operation; - let progress = download.progress.clone(); - if let Some(result) = result { - match result { - Ok(DownloadOutcome::Complete) => { - let progress = model::artifact_download_progress(artifact, &models_path()); - self.model_download = ModelDownload::Complete(artifact, operation, progress); - self.error = None; - } - Ok(DownloadOutcome::Stopped) => { - self.model_download = ModelDownload::Idle; - self.error = None; - } - Err(error) => { - self.error = Some(error.clone()); - self.model_download = ModelDownload::Failed(artifact, error, progress); - } - } - } - } - - fn start_generation(&mut self) { - if self.generating || self.selected_session.is_none() { - return; - } - let prompt = self.composer.trim().to_owned(); - if prompt.is_empty() { - return; - } - let model = match ModelChoice::from_id(&self.preferences.selected_model) { - Some(model) => model, - None => { - self.error = Some("The selected model is not supported.".into()); - return; - } - }; - let effective = self.preferences.generation().and_then(|generation| { - self.preferences.runtime().and_then(|runtime| { - crate::settings::effective_settings(model, &generation, &runtime, &models_path()) - }) - }); - let effective = match effective { - Ok(settings) => settings, - Err(error) => { - self.error = Some(error); - return; - } - }; - let assistant_reasoning = effective.turn.reasoning_mode != ReasoningMode::Direct; - let session_id = self - .selected_session - .expect("a selected session was checked"); - #[cfg(target_os = "macos")] - let mut messages = self - .conversation - .iter() - .map(|message| ChatTurn { - user: message.user, - skip_previous_eos: false, - reasoning: message.reasoning.clone(), - reasoning_complete: message.reasoning_complete, - content: message.content.clone(), - }) - .collect::>(); - #[cfg(target_os = "macos")] - messages.push(ChatTurn { - user: true, - skip_previous_eos: false, - reasoning: None, - reasoning_complete: true, - content: prompt.clone(), - }); - - #[cfg(target_os = "macos")] - { - let Some(service) = &self.generation_service else { - self.error = Some("The model runtime is unavailable.".into()); - return; - }; - let Some(database) = &mut self.database else { - return; - }; - let saved = match database.start_chat_turn(session_id, &prompt, assistant_reasoning) { - Ok(turn) => turn, - Err(error) => { - self.error = Some(format!("Could not save the chat turn: {error}")); - return; - } - }; - let idle_timeout = - Duration::from_secs(self.preferences.idle_timeout_minutes.max(1) as u64 * 60); - self.active_generation = match service.generate( - effective.engine, - effective.turn, - messages, - CheckpointTarget::Local(session_checkpoint_path(session_id)), - idle_timeout, - ) { - Ok(active) => Some(active), - Err(error) => { - self.generation_service = None; - self.error = Some(error); - return; - } - }; - let user = ChatMessage::from(saved.0); - let mut assistant = ChatMessage::from(saved.1); - assistant.reasoning_open = assistant_reasoning; - self.composer.clear(); - self.conversation.push(user); - self.conversation.push(assistant); - self.generating = true; - self.tokens_per_second = None; - self.error = None; - } - #[cfg(not(target_os = "macos"))] - { - let _ = effective; - self.error = Some("Local Metal generation requires macOS.".into()); - return; - } - } - - fn poll_generation(&mut self) -> bool { - #[cfg(target_os = "macos")] - let Some(active) = &mut self.active_generation else { - self.generating = false; - return false; - }; - #[cfg(target_os = "macos")] - let mut transcript_changed = false; - #[cfg(target_os = "macos")] - let mut context_changed = false; - #[cfg(target_os = "macos")] - loop { - match active.events.try_recv() { - Ok(GenerationEvent::Loading) => {} - Ok(GenerationEvent::Chunk { reasoning, content }) => { - if let Some(message) = self.conversation.last_mut() - && !message.user - { - message.append(reasoning, &content); - transcript_changed = true; - } - } - Ok(GenerationEvent::Context { - used, - limit, - tokens_per_second, - }) => { - self.context_used = used; - self.context_limit = limit; - self.tokens_per_second = tokens_per_second; - context_changed = true; - } - Ok(GenerationEvent::Finished(result)) => { - self.generating = false; - if let Err(error) = result { - self.error = Some(error); - } - self.active_generation = None; - break; - } - Err(TryRecvError::Empty) => break, - Err(TryRecvError::Disconnected) => { - self.generating = false; - self.active_generation = None; - self.error = Some("The model runtime stopped unexpectedly.".into()); - break; - } - } - } - #[cfg(target_os = "macos")] - if transcript_changed && let Some(message) = self.conversation.last_mut() { - message.refresh_markdown(); - } - #[cfg(target_os = "macos")] - if transcript_changed - && let Some(message) = self.conversation.last() - && let Some(database) = &mut self.database - && let Err(error) = database.update_message( - message.id, - message.reasoning.as_deref(), - message.reasoning_complete, - &message.content, - ) - { - if let Some(active) = &self.active_generation { - active.cancel.store(true, Ordering::Relaxed); - } - self.error = Some(format!("Could not save generated chat text: {error}")); - } - #[cfg(target_os = "macos")] - if context_changed - && let Some(session_id) = self.selected_session - && let Some(database) = &mut self.database - { - if let Err(error) = database.update_session_context( - session_id, - self.context_used, - self.context_limit, - self.tokens_per_second, - ) { - self.error = Some(format!("Could not save context usage: {error}")); - } else if let Some(session) = self - .projects - .iter_mut() - .flat_map(|project| &mut project.sessions) - .find(|session| session.id == session_id) - { - session.context_used = self.context_used as i32; - session.context_limit = self.context_limit as i32; - session.last_tokens_per_second = self.tokens_per_second; - } - } - #[cfg(target_os = "macos")] - return transcript_changed; - #[cfg(not(target_os = "macos"))] - false - } - fn sample_metrics(&mut self) { let prefill_tokens_per_second = self.metrics.take_prefill_sample(); let (kv_read_bytes, kv_write_bytes) = self.metrics.take_kv_io_sample(); diff --git a/src/app/generation.rs b/src/app/generation.rs new file mode 100644 index 0000000..fec2bf1 --- /dev/null +++ b/src/app/generation.rs @@ -0,0 +1,252 @@ +use super::*; + +#[derive(Clone, Debug)] +pub(crate) struct ChatMessage { + pub(super) id: i32, + pub(super) user: bool, + pub(super) reasoning: Option, + pub(super) reasoning_complete: bool, + pub(super) reasoning_open: bool, + pub(super) content: String, + pub(super) markdown: Vec, +} + +impl ChatMessage { + pub(super) fn append(&mut self, reasoning: bool, chunk: &str) { + if reasoning { + self.reasoning.get_or_insert_default().push_str(chunk); + } else { + self.reasoning_complete = true; + self.content.push_str(chunk); + } + } + + pub(super) fn refresh_markdown(&mut self) { + if !self.user { + let content = if self.reasoning.is_some() { + self.content.trim_start() + } else { + &self.content + }; + self.markdown = markdown::parse(content).collect(); + } + } +} + +impl From for ChatMessage { + fn from(message: StoredMessage) -> Self { + let mut message = Self { + id: message.id, + user: message.user, + reasoning: message.reasoning, + reasoning_complete: message.reasoning_complete, + reasoning_open: false, + content: message.content, + markdown: Vec::new(), + }; + message.refresh_markdown(); + message + } +} + +impl App { + pub(super) fn start_generation(&mut self) { + if self.generating || self.selected_session.is_none() { + return; + } + let prompt = self.composer.trim().to_owned(); + if prompt.is_empty() { + return; + } + let model = match ModelChoice::from_id(&self.preferences.selected_model) { + Some(model) => model, + None => { + self.error = Some("The selected model is not supported.".into()); + return; + } + }; + let effective = self.preferences.generation().and_then(|generation| { + self.preferences.runtime().and_then(|runtime| { + crate::settings::effective_settings(model, &generation, &runtime, &models_path()) + }) + }); + let effective = match effective { + Ok(settings) => settings, + Err(error) => { + self.error = Some(error); + return; + } + }; + let assistant_reasoning = effective.turn.reasoning_mode != ReasoningMode::Direct; + let session_id = self + .selected_session + .expect("a selected session was checked"); + #[cfg(target_os = "macos")] + let mut messages = self + .conversation + .iter() + .map(|message| ChatTurn { + user: message.user, + skip_previous_eos: false, + reasoning: message.reasoning.clone(), + reasoning_complete: message.reasoning_complete, + content: message.content.clone(), + }) + .collect::>(); + #[cfg(target_os = "macos")] + messages.push(ChatTurn { + user: true, + skip_previous_eos: false, + reasoning: None, + reasoning_complete: true, + content: prompt.clone(), + }); + + #[cfg(target_os = "macos")] + { + let Some(service) = &self.generation_service else { + self.error = Some("The model runtime is unavailable.".into()); + return; + }; + let Some(database) = &mut self.database else { + return; + }; + let saved = match database.start_chat_turn(session_id, &prompt, assistant_reasoning) { + Ok(turn) => turn, + Err(error) => { + self.error = Some(format!("Could not save the chat turn: {error}")); + return; + } + }; + let idle_timeout = + Duration::from_secs(self.preferences.idle_timeout_minutes.max(1) as u64 * 60); + self.active_generation = match service.generate( + effective.engine, + effective.turn, + messages, + CheckpointTarget::Local(session_checkpoint_path(session_id)), + idle_timeout, + ) { + Ok(active) => Some(active), + Err(error) => { + self.generation_service = None; + self.error = Some(error); + return; + } + }; + let user = ChatMessage::from(saved.0); + let mut assistant = ChatMessage::from(saved.1); + assistant.reasoning_open = assistant_reasoning; + self.composer.clear(); + self.conversation.push(user); + self.conversation.push(assistant); + self.generating = true; + self.tokens_per_second = None; + self.error = None; + } + #[cfg(not(target_os = "macos"))] + { + let _ = effective; + self.error = Some("Local Metal generation requires macOS.".into()); + return; + } + } + + pub(super) fn poll_generation(&mut self) -> bool { + #[cfg(target_os = "macos")] + let Some(active) = &mut self.active_generation else { + self.generating = false; + return false; + }; + #[cfg(target_os = "macos")] + let mut transcript_changed = false; + #[cfg(target_os = "macos")] + let mut context_changed = false; + #[cfg(target_os = "macos")] + loop { + match active.events.try_recv() { + Ok(GenerationEvent::Loading) => {} + Ok(GenerationEvent::Chunk { reasoning, content }) => { + if let Some(message) = self.conversation.last_mut() + && !message.user + { + message.append(reasoning, &content); + transcript_changed = true; + } + } + Ok(GenerationEvent::Context { + used, + limit, + tokens_per_second, + }) => { + self.context_used = used; + self.context_limit = limit; + self.tokens_per_second = tokens_per_second; + context_changed = true; + } + Ok(GenerationEvent::Finished(result)) => { + self.generating = false; + if let Err(error) = result { + self.error = Some(error); + } + self.active_generation = None; + break; + } + Err(TryRecvError::Empty) => break, + Err(TryRecvError::Disconnected) => { + self.generating = false; + self.active_generation = None; + self.error = Some("The model runtime stopped unexpectedly.".into()); + break; + } + } + } + #[cfg(target_os = "macos")] + if transcript_changed && let Some(message) = self.conversation.last_mut() { + message.refresh_markdown(); + } + #[cfg(target_os = "macos")] + if transcript_changed + && let Some(message) = self.conversation.last() + && let Some(database) = &mut self.database + && let Err(error) = database.update_message( + message.id, + message.reasoning.as_deref(), + message.reasoning_complete, + &message.content, + ) + { + if let Some(active) = &self.active_generation { + active.cancel.store(true, Ordering::Relaxed); + } + self.error = Some(format!("Could not save generated chat text: {error}")); + } + #[cfg(target_os = "macos")] + if context_changed + && let Some(session_id) = self.selected_session + && let Some(database) = &mut self.database + { + if let Err(error) = database.update_session_context( + session_id, + self.context_used, + self.context_limit, + self.tokens_per_second, + ) { + self.error = Some(format!("Could not save context usage: {error}")); + } else if let Some(session) = self + .projects + .iter_mut() + .flat_map(|project| &mut project.sessions) + .find(|session| session.id == session_id) + { + session.context_used = self.context_used as i32; + session.context_limit = self.context_limit as i32; + session.last_tokens_per_second = self.tokens_per_second; + } + } + #[cfg(target_os = "macos")] + return transcript_changed; + #[cfg(not(target_os = "macos"))] + false + } +} diff --git a/src/app/model_manager.rs b/src/app/model_manager.rs new file mode 100644 index 0000000..72d3fef --- /dev/null +++ b/src/app/model_manager.rs @@ -0,0 +1,173 @@ +use super::*; + +#[derive(Debug)] +pub(super) enum ModelDownload { + Idle, + Active(ActiveDownload), + Complete(ManagedArtifactId, ModelOperation, DownloadProgress), + Failed(ManagedArtifactId, String, DownloadProgress), +} + +#[derive(Clone, Copy, Debug, Eq, PartialEq)] +pub(super) enum ModelOperation { + Download, + Validate, +} + +#[derive(Debug)] +pub(super) struct ActiveDownload { + pub(super) artifact: ManagedArtifactId, + pub(super) operation: ModelOperation, + pub(super) progress: DownloadProgress, + pub(super) sampled_at: Instant, + pub(super) sampled_bytes: u64, + pub(super) bytes_per_second: f64, + pub(super) verified_bytes: Arc, + pub(super) cancel: Arc, + pub(super) result: mpsc::Receiver>, + pub(super) stopping: bool, +} + +impl App { + pub(super) fn open_model_manager(&mut self) -> Task { + if let Some(id) = self.model_manager_window { + return window::gain_focus(id); + } + let (id, open) = window::open(window::Settings { + size: Size::new(760.0, 560.0), + min_size: Some(Size::new(620.0, 420.0)), + icon: Some(app_icon()), + ..Default::default() + }); + self.model_manager_window = Some(id); + open.map(Message::ModelManagerOpened) + } + + pub(super) fn start_model_operation( + &mut self, + artifact: ManagedArtifactId, + operation: ModelOperation, + ) { + if matches!(self.model_download, ModelDownload::Active(_)) { + return; + } + let models_path = models_path(); + let progress = match operation { + ModelOperation::Download => model::artifact_download_progress(artifact, &models_path), + ModelOperation::Validate => model::artifact_verification_progress(artifact, 0), + }; + let cancel = Arc::new(AtomicBool::new(false)); + let worker_cancel = Arc::clone(&cancel); + let verified_bytes = Arc::new(AtomicU64::new(0)); + let worker_verified_bytes = Arc::clone(&verified_bytes); + let (result_sender, result_receiver) = mpsc::channel(); + let thread_name = match operation { + ModelOperation::Download => "model-download", + ModelOperation::Validate => "model-validation", + }; + if let Err(error) = thread::Builder::new() + .name(thread_name.to_owned()) + .spawn(move || { + let result = match operation { + ModelOperation::Download => model::download_managed_artifact( + artifact, + &models_path, + &worker_cancel, + &worker_verified_bytes, + ), + ModelOperation::Validate => model::validate_managed_artifact( + artifact, + &models_path, + &worker_cancel, + &worker_verified_bytes, + ), + }; + let _ = result_sender.send(result); + }) + { + self.error = Some(format!("Could not start {thread_name}: {error}")); + return; + } + self.pending_model_delete = None; + self.model_download = ModelDownload::Active(ActiveDownload { + artifact, + operation, + sampled_at: Instant::now(), + sampled_bytes: progress.completed(), + progress, + bytes_per_second: 0.0, + verified_bytes, + cancel, + result: result_receiver, + stopping: false, + }); + } +} + +impl App { + pub(super) fn update_download_progress(&mut self) { + let ModelDownload::Active(download) = &mut self.model_download else { + return; + }; + let now = Instant::now(); + let verified = download.verified_bytes.load(Ordering::Relaxed); + let mut progress = match download.operation { + ModelOperation::Download => { + model::artifact_download_progress(download.artifact, &models_path()) + } + ModelOperation::Validate => { + model::artifact_verification_progress(download.artifact, verified) + } + }; + if download.operation == ModelOperation::Download + && let Some(verification) = &mut progress.verification + { + verification.verified = verified.min(verification.total); + } + let elapsed = now.duration_since(download.sampled_at).as_secs_f64(); + let completed = progress.completed(); + let phase_changed = progress.phase != download.progress.phase; + let transferred = completed.saturating_sub(download.sampled_bytes); + if phase_changed { + download.bytes_per_second = 0.0; + } else if transferred > 0 && elapsed > 0.0 { + let current = transferred as f64 / elapsed; + download.bytes_per_second = if download.bytes_per_second == 0.0 { + current + } else { + download.bytes_per_second * 0.75 + current * 0.25 + }; + } + download.progress = progress; + download.sampled_at = now; + download.sampled_bytes = completed; + + let result = match download.result.try_recv() { + Ok(result) => Some(result), + Err(TryRecvError::Empty) => None, + Err(TryRecvError::Disconnected) => { + Some(Err("Download worker stopped unexpectedly.".into())) + } + }; + let artifact = download.artifact; + let operation = download.operation; + let progress = download.progress.clone(); + if let Some(result) = result { + match result { + Ok(DownloadOutcome::Complete) => { + let progress = model::artifact_download_progress(artifact, &models_path()); + self.model_download = ModelDownload::Complete(artifact, operation, progress); + self.error = None; + } + Ok(DownloadOutcome::Stopped) => { + self.model_download = ModelDownload::Idle; + self.error = None; + } + Err(error) => { + self.error = Some(error.clone()); + self.model_download = ModelDownload::Failed(artifact, error, progress); + } + } + } + } +} diff --git a/src/app/preferences.rs b/src/app/preferences.rs new file mode 100644 index 0000000..0c307ce --- /dev/null +++ b/src/app/preferences.rs @@ -0,0 +1,476 @@ +use super::*; + +#[derive(Clone)] +pub(super) struct PreferenceDraft { + pub(super) model: ModelChoice, + pub(super) dspark_enabled: bool, + pub(super) idle_timeout_minutes: String, + pub(super) endpoint_port: String, + pub(super) context_tokens: String, + pub(super) max_generated_tokens: String, + pub(super) system_prompt: String, + pub(super) temperature: String, + pub(super) top_p: String, + pub(super) min_p: String, + pub(super) seed: String, + pub(super) reasoning_mode: ReasoningMode, + pub(super) cpu_threads: String, + pub(super) power_percent: String, + pub(super) prefill_chunk: String, + pub(super) quality: bool, + pub(super) warm_weights: bool, + pub(super) mtp_draft_tokens: String, + pub(super) mtp_margin: String, + pub(super) glm_mtp: bool, + pub(super) glm_mtp_timing: bool, + pub(super) dspark_confidence_threshold: String, + pub(super) dspark_strict: bool, + pub(super) ssd_streaming: bool, + pub(super) ssd_streaming_cold: bool, + pub(super) ssd_cache: String, + pub(super) ssd_full_layers: String, + pub(super) ssd_preload_experts: String, + pub(super) directional_steering_file: String, + pub(super) directional_steering_ffn: String, + pub(super) directional_steering_attn: String, + pub(super) simulated_used_memory_gib: String, + pub(super) expert_profile_path: String, +} + +impl PreferenceDraft { + pub(super) fn from_saved(preferences: &AppPreferences) -> Result { + let model = ModelChoice::from_id(&preferences.selected_model) + .ok_or_else(|| format!("Unsupported model: {}", preferences.selected_model))?; + let generation = preferences.generation()?; + let runtime = preferences.runtime()?; + runtime.validate(model)?; + let execution = &runtime.execution; + let speculative = &runtime.speculative; + Ok(Self { + model, + dspark_enabled: speculative.dspark_enabled, + idle_timeout_minutes: preferences.idle_timeout_minutes.to_string(), + endpoint_port: preferences.endpoint_port.to_string(), + context_tokens: generation.context_tokens.to_string(), + max_generated_tokens: generation.max_generated_tokens.to_string(), + system_prompt: generation.system_prompt, + temperature: generation + .temperature + .map_or_else(String::new, |value| value.to_string()), + top_p: generation + .top_p + .map_or_else(String::new, |value| value.to_string()), + min_p: generation + .min_p + .map_or_else(String::new, |value| value.to_string()), + seed: generation + .seed + .map_or_else(String::new, |value| value.to_string()), + reasoning_mode: generation.reasoning_mode, + cpu_threads: execution + .cpu_threads + .map_or_else(String::new, |value| value.to_string()), + power_percent: execution + .power_percent + .map_or_else(String::new, |value| value.to_string()), + prefill_chunk: execution + .prefill_chunk + .map_or_else(String::new, |value| value.to_string()), + quality: execution.quality, + warm_weights: execution.warm_weights, + mtp_draft_tokens: speculative.mtp_draft_tokens.to_string(), + mtp_margin: speculative.mtp_margin.to_string(), + glm_mtp: speculative.glm_mtp, + glm_mtp_timing: speculative.glm_mtp_timing, + dspark_confidence_threshold: speculative + .dspark_confidence_threshold + .map_or_else(String::new, |value| value.to_string()), + dspark_strict: speculative.dspark_strict, + ssd_streaming: runtime.ssd.enabled, + ssd_streaming_cold: runtime.ssd.cold, + ssd_cache: runtime + .ssd + .cache + .map_or_else(String::new, |cache| match cache { + StreamingCacheBudget::Experts(experts) => experts.to_string(), + StreamingCacheBudget::Gib(gib) => format!("{gib}GB"), + }), + ssd_full_layers: runtime + .ssd + .full_layers + .map_or_else(String::new, |value| value.to_string()), + ssd_preload_experts: runtime + .ssd + .preload_experts + .map_or_else(String::new, |value| value.to_string()), + directional_steering_file: runtime.steering.file.unwrap_or_default(), + directional_steering_ffn: runtime + .steering + .ffn_scale + .map_or_else(String::new, |value| value.to_string()), + directional_steering_attn: runtime + .steering + .attention_scale + .map_or_else(String::new, |value| value.to_string()), + simulated_used_memory_gib: runtime + .diagnostics + .simulated_used_memory_gib + .map_or_else(String::new, |value| value.to_string()), + expert_profile_path: runtime.diagnostics.expert_profile_path.unwrap_or_default(), + }) + } + + pub(super) fn generation(&self) -> Result { + let preferences = GenerationPreferences { + context_tokens: parse_positive_i32("Context tokens", &self.context_tokens)?, + max_generated_tokens: parse_positive_i32( + "Maximum generated tokens", + &self.max_generated_tokens, + )?, + system_prompt: self.system_prompt.clone(), + temperature: parse_optional_f32("Temperature", &self.temperature)?, + top_p: parse_optional_f32("Top-p", &self.top_p)?, + min_p: parse_optional_f32("Min-p", &self.min_p)?, + seed: parse_optional_u64("Seed", &self.seed)?, + reasoning_mode: self.reasoning_mode, + }; + preferences.validate()?; + Ok(preferences) + } + + pub(super) fn reset(&mut self) { + let defaults = GenerationPreferences::default(); + let execution = ExecutionPreferences::default(); + let speculative = SpeculativePreferences::default(); + let ssd = SsdPreferences::default(); + self.model = ModelChoice::default(); + self.dspark_enabled = false; + self.idle_timeout_minutes = "10".into(); + self.endpoint_port = "4000".into(); + self.context_tokens = defaults.context_tokens.to_string(); + self.max_generated_tokens = defaults.max_generated_tokens.to_string(); + self.system_prompt = defaults.system_prompt; + self.temperature.clear(); + self.top_p.clear(); + self.min_p.clear(); + self.seed.clear(); + self.reasoning_mode = defaults.reasoning_mode; + self.cpu_threads.clear(); + self.power_percent.clear(); + self.prefill_chunk.clear(); + self.quality = execution.quality; + self.warm_weights = execution.warm_weights; + self.mtp_draft_tokens = speculative.mtp_draft_tokens.to_string(); + self.mtp_margin = speculative.mtp_margin.to_string(); + self.glm_mtp = speculative.glm_mtp; + self.glm_mtp_timing = speculative.glm_mtp_timing; + self.dspark_confidence_threshold.clear(); + self.dspark_strict = speculative.dspark_strict; + self.ssd_streaming = ssd.enabled; + self.ssd_streaming_cold = ssd.cold; + self.ssd_cache.clear(); + self.ssd_full_layers.clear(); + self.ssd_preload_experts.clear(); + self.directional_steering_file.clear(); + self.directional_steering_ffn.clear(); + self.directional_steering_attn.clear(); + self.simulated_used_memory_gib.clear(); + self.expert_profile_path.clear(); + } + + pub(super) fn execution(&self) -> Result { + Ok(ExecutionPreferences { + cpu_threads: parse_optional_u32("CPU helper threads", &self.cpu_threads)?, + power_percent: parse_optional_u8("GPU power", &self.power_percent)?, + prefill_chunk: parse_optional_u32("Prefill chunk", &self.prefill_chunk)?, + quality: self.quality, + warm_weights: self.warm_weights, + }) + } + + pub(super) fn speculative(&self) -> Result { + Ok(SpeculativePreferences { + mtp_draft_tokens: parse_positive_i32("MTP draft tokens", &self.mtp_draft_tokens)?, + mtp_margin: parse_f32("MTP margin", &self.mtp_margin)?, + glm_mtp: self.glm_mtp, + glm_mtp_timing: self.glm_mtp_timing, + dspark_enabled: self.dspark_enabled, + dspark_confidence_threshold: parse_optional_f32( + "DSpark confidence", + &self.dspark_confidence_threshold, + )?, + dspark_strict: self.dspark_strict, + }) + } + + pub(super) fn runtime(&self) -> Result { + Ok(RuntimePreferences { + execution: self.execution()?, + speculative: self.speculative()?, + ssd: SsdPreferences { + enabled: self.ssd_streaming, + cold: self.ssd_streaming_cold, + cache: parse_streaming_cache(&self.ssd_cache)?, + full_layers: parse_optional_u32("SSD full-layer count", &self.ssd_full_layers)?, + preload_experts: parse_optional_u32( + "SSD preload experts", + &self.ssd_preload_experts, + )?, + }, + steering: SteeringPreferences { + file: optional_text(&self.directional_steering_file), + ffn_scale: parse_optional_f32( + "Directional FFN scale", + &self.directional_steering_ffn, + )?, + attention_scale: parse_optional_f32( + "Directional attention scale", + &self.directional_steering_attn, + )?, + }, + diagnostics: DiagnosticPreferences { + simulated_used_memory_gib: parse_optional_gib( + "Simulated used memory", + &self.simulated_used_memory_gib, + )?, + expert_profile_path: optional_text(&self.expert_profile_path), + }, + }) + } +} + +fn parse_positive_i32(name: &str, value: &str) -> Result { + value + .trim() + .parse::() + .ok() + .filter(|value| *value > 0) + .ok_or_else(|| format!("{name} must be a positive whole number.")) +} + +fn parse_optional_f32(name: &str, value: &str) -> Result, String> { + let value = value.trim(); + if value.is_empty() { + Ok(None) + } else { + value + .parse() + .map(Some) + .map_err(|_| format!("{name} must be a number or left blank for the DS4 default.")) + } +} + +fn parse_f32(name: &str, value: &str) -> Result { + value + .trim() + .parse() + .map_err(|_| format!("{name} must be a number.")) +} + +fn parse_optional_u64(name: &str, value: &str) -> Result, String> { + let value = value.trim(); + if value.is_empty() { + Ok(None) + } else { + value + .parse() + .map(Some) + .map_err(|_| format!("{name} must be a positive whole number or left blank.")) + } +} + +fn parse_optional_u32(name: &str, value: &str) -> Result, String> { + parse_optional_number(name, value) +} + +fn parse_optional_u8(name: &str, value: &str) -> Result, String> { + parse_optional_number(name, value) +} + +fn parse_optional_number( + name: &str, + value: &str, +) -> Result, String> { + let value = value.trim(); + if value.is_empty() { + Ok(None) + } else { + value + .parse() + .map(Some) + .map_err(|_| format!("{name} must be a positive whole number or left blank.")) + } +} + +pub(super) fn parse_streaming_cache(value: &str) -> Result, String> { + let value = value.trim(); + if value.is_empty() { + return Ok(None); + } + if value.len() > 2 + && value + .get(value.len() - 2..) + .is_some_and(|suffix| suffix.eq_ignore_ascii_case("gb")) + { + return parse_gib("SSD cache budget", value) + .map(|gib| Some(StreamingCacheBudget::Gib(gib))); + } + if !value.chars().all(|character| character.is_ascii_digit()) { + return Err( + "SSD cache budget must be a positive expert count or whole GiB value such as 64GB." + .into(), + ); + } + value + .parse::() + .ok() + .filter(|value| *value > 0) + .map(StreamingCacheBudget::Experts) + .map(Some) + .ok_or_else(|| { + "SSD cache budget must be a positive expert count or whole GiB value such as 64GB." + .into() + }) +} + +pub(super) fn parse_optional_gib(name: &str, value: &str) -> Result, String> { + let value = value.trim(); + if value.is_empty() { + Ok(None) + } else { + parse_gib(name, value).map(Some) + } +} + +fn parse_gib(name: &str, value: &str) -> Result { + let value = value + .get(value.len().saturating_sub(2)..) + .filter(|suffix| suffix.eq_ignore_ascii_case("gb")) + .map_or(value, |_| &value[..value.len() - 2]); + if !value.chars().all(|character| character.is_ascii_digit()) { + return Err(format!("{name} must be a positive whole GiB value.")); + } + value + .parse::() + .ok() + .filter(|value| *value > 0 && *value <= u64::MAX / GIB) + .ok_or_else(|| format!("{name} must be a positive whole GiB value.")) +} + +fn optional_text(value: &str) -> Option { + let value = value.trim(); + (!value.is_empty()).then(|| value.to_owned()) +} + +impl App { + pub(super) fn open_preferences(&mut self) { + if self.database.is_none() || self.pending_project_path.is_some() || self.choosing_folder { + return; + } + match PreferenceDraft::from_saved(&self.preferences) { + Ok(draft) => { + self.preference_draft = draft; + self.preference_error = None; + self.preferences_open = true; + } + Err(error) => self.error = Some(error), + } + } + + pub(super) fn save_preferences(&mut self) { + let Ok(idle_timeout_minutes) = self + .preference_draft + .idle_timeout_minutes + .trim() + .parse::() + else { + self.preference_error = Some("Idle timeout must be a whole number.".into()); + return; + }; + if !(1..=1440).contains(&idle_timeout_minutes) { + self.preference_error = Some("Idle timeout must be between 1 and 1440 minutes.".into()); + return; + } + let Ok(endpoint_port) = self.preference_draft.endpoint_port.trim().parse::() else { + self.preference_error = Some("Endpoint port must be a whole number.".into()); + return; + }; + if endpoint_port == 0 { + self.preference_error = Some("Endpoint port must be between 1 and 65535.".into()); + return; + } + let generation = match self.preference_draft.generation() { + Ok(generation) => generation, + Err(error) => { + self.preference_error = Some(error); + return; + } + }; + + let model = self.preference_draft.model; + let runtime = match self.preference_draft.runtime() { + Ok(runtime) => runtime, + Err(error) => { + self.preference_error = Some(error); + return; + } + }; + if let Err(error) = runtime.validate(model) { + self.preference_error = Some(error); + return; + } + #[cfg(target_os = "macos")] + let pending_endpoint = if self.preferences.endpoint_port != i32::from(endpoint_port) + || self._endpoint.is_none() + { + let Some(generation) = &self.generation_service else { + self.preference_error = Some("The model runtime is unavailable.".into()); + return; + }; + match crate::server::ServerHandle::spawn( + generation.clone(), + Arc::clone(&self.runtime_preferences), + models_path(), + application_support_path().join("kv-cache").join("http"), + endpoint_port, + Arc::clone(&self.metrics), + ) { + Ok(endpoint) => Some(endpoint), + Err(error) => { + self.preference_error = Some(error); + return; + } + } + } else { + None + }; + let Some(database) = &mut self.database else { + return; + }; + match database.update_preferences( + model.id(), + idle_timeout_minutes, + i32::from(endpoint_port), + &generation, + &runtime, + ) { + Ok(preferences) => { + self.preferences = preferences; + #[cfg(target_os = "macos")] + if let Ok(mut runtime_preferences) = self.runtime_preferences.write() { + *runtime_preferences = self.preferences.clone(); + } + #[cfg(target_os = "macos")] + if let Some(endpoint) = pending_endpoint { + self._endpoint = Some(endpoint); + } + self.preference_draft = PreferenceDraft::from_saved(&self.preferences) + .expect("the saved model was selected from the supported catalog"); + self.preferences_open = false; + self.preference_error = None; + self.error = None; + } + Err(error) => self.preference_error = Some(error), + } + } +} diff --git a/src/app/projects.rs b/src/app/projects.rs new file mode 100644 index 0000000..a2ae827 --- /dev/null +++ b/src/app/projects.rs @@ -0,0 +1,102 @@ +use super::*; + +impl App { + pub(super) fn prepare_project(&mut self, path: PathBuf) { + let Ok(path) = fs::canonicalize(path) else { + self.error = Some("The selected folder is no longer available.".into()); + return; + }; + let Some(path_text) = path.to_str() else { + self.error = Some("The selected folder path is not valid UTF-8.".into()); + return; + }; + if self + .projects + .iter() + .any(|item| item.project.path == path_text) + { + self.error = Some("That project is already in the sidebar.".into()); + return; + } + + self.project_name_input = path + .file_name() + .and_then(|name| name.to_str()) + .unwrap_or("Project") + .to_owned(); + self.pending_project_path = Some(path); + self.error = None; + } + + pub(super) fn create_project(&mut self) { + let name = self.project_name_input.trim(); + if name.is_empty() { + self.error = Some("Project name cannot be empty.".into()); + return; + } + let Some(path) = &self.pending_project_path else { + return; + }; + let Some(path) = path.to_str() else { + self.error = Some("The selected folder path is not valid UTF-8.".into()); + return; + }; + let Some(database) = &mut self.database else { + return; + }; + + match database.create_project(name, path) { + Ok(project) => { + self.selected_project = Some(project.id); + self.selected_session = None; + self.pending_project_path = None; + self.project_name_input.clear(); + self.error = None; + self.reload_projects(); + } + Err(error) => self.error = Some(error), + } + } + + pub(super) fn create_session(&mut self) { + if self.generating { + self.error = Some("Stop the active generation before creating a session.".into()); + return; + } + let Some(project_id) = self.selected_project else { + self.error = Some("Select a project first.".into()); + return; + }; + let default_number = self + .selected_project() + .map(|project| project.sessions.len() + 1) + .unwrap_or(1); + let title = format!("Session {default_number}"); + let Some(database) = &mut self.database else { + return; + }; + + match database.create_session(project_id, &title) { + Ok(session) => { + self.selected_session = Some(session.id); + self.conversation.clear(); + self.composer.clear(); + self.context_used = 0; + self.context_limit = self.preferences.context_tokens.max(0) as u32; + self.tokens_per_second = None; + self.error = None; + self.reload_projects(); + } + Err(error) => self.error = Some(error), + } + } + + pub(super) fn reload_projects(&mut self) { + if let Some(database) = &mut self.database { + match database.load_projects() { + Ok(projects) => self.projects = projects, + Err(error) => self.error = Some(error), + } + } + } +} diff --git a/src/app/view.rs b/src/app/view.rs index 5bc5d90..324164b 100644 --- a/src/app/view.rs +++ b/src/app/view.rs @@ -1,3 +1,10 @@ +mod chat; +mod model_manager; +mod preferences; +mod stats; + +use model_manager::{download_status_bar, format_bytes, format_duration}; + use super::{ ActiveDownload, App, DetailTab, Message, MetricsPoint, ModelDownload, ModelOperation, chat_scroll_id, models_path, @@ -229,1088 +236,6 @@ impl App { .into() } - fn chat_detail(&self) -> Element<'_, Message> { - let Some(item) = self.selected_project() else { - let open_project_content = row![icon(ICON_FOLDER_PLUS, 17), text("Open project…"),] - .spacing(8) - .align_y(Alignment::Center); - let open_project = if self.database.is_some() && !self.choosing_folder { - action_button(open_project_content).on_press(Message::ChooseProjectFolder) - } else { - action_button(open_project_content) - }; - return container( - column![ - icon(ICON_SPARK, 36), - text("Start a local coding session").size(28), - text("Choose a project folder to create your first session.").size(14), - Space::with_height(10), - open_project, - ] - .spacing(10) - .align_x(Alignment::Center), - ) - .center_x(Length::Fill) - .center_y(Length::Fill) - .into(); - }; - - let project = &item.project; - let selected_title = self - .selected_session(item) - .map(|session| session.title.as_str()) - .unwrap_or(project.name.as_str()); - let header = row![ - icon(ICON_FOLDER, 19), - text(selected_title).size(18), - icon(ICON_MORE, 18), - Space::with_width(Length::Fill), - ] - .spacing(10) - .align_y(Alignment::Center); - - let body: Element<'_, Message> = if let Some(session) = self.selected_session(item) { - let mut messages = column![].spacing(12); - if self.conversation.is_empty() { - messages = messages.push( - column![ - text(&session.title).size(26), - text("Run DeepSeek locally with the Rust Metal engine.").size(14), - ] - .spacing(8), - ); - } else { - let markdown_style = markdown::Style::from_palette(app_theme().palette()); - for (index, message) in self.conversation.iter().enumerate() { - let label = if message.user { "You" } else { "DS4" }; - let active = self.generating && index + 1 == self.conversation.len(); - let mut body = column![text(label).size(11)].spacing(5); - if let Some(reasoning) = &message.reasoning { - let reasoning_label = - match (message.reasoning_open, message.reasoning_complete, active) { - (true, false, true) => "▾ Thinking", - (false, false, true) => "› Thinking", - (true, false, false) => "▾ Reasoning (stopped)", - (false, false, false) => "› Reasoning (stopped)", - (true, true, _) => "▾ Reasoning", - (false, true, _) => "› Reasoning", - }; - body = body.push( - button(text(reasoning_label).size(12)) - .padding(0) - .style(button::text) - .on_press(Message::ToggleReasoning(index)), - ); - if message.reasoning_open { - body = body.push( - text(if reasoning.is_empty() && active { - "Thinking…" - } else { - reasoning - }) - .size(13) - .color(muted_text()), - ); - } - } - if !message.content.is_empty() { - if message.user || message.markdown.is_empty() { - let content = if message.reasoning.is_some() { - message.content.trim_start() - } else { - &message.content - }; - body = body.push(text(content).size(14)); - } else { - body = body.push( - markdown::view( - &message.markdown, - markdown::Settings::with_text_size(14), - markdown_style, - ) - .map(Message::OpenLink), - ); - } - } else if active && message.reasoning.is_none() { - body = body.push(text("Loading model…").size(14)); - } - let user = message.user; - messages = messages.push( - container(body) - .padding(14) - .width(Length::Fill) - .style(move |theme| chat_message_style(theme, user)), - ); - } - } - let composer = text_input("Ask DS4Server anything…", &self.composer) - .on_input(Message::ComposerChanged) - .on_submit(Message::SubmitPrompt) - .padding(12) - .size(14); - let action = if self.generating { - action_button(text("Stop").size(12)).on_press(Message::StopGeneration) - } else if self.composer.trim().is_empty() { - action_button(icon(ICON_SEND, 18)).padding(8) - } else { - action_button(icon(ICON_SEND, 18)) - .padding(8) - .on_press(Message::SubmitPrompt) - }; - let context_fraction = if self.context_limit == 0 { - 0.0 - } else { - self.context_used.min(self.context_limit) as f32 / self.context_limit as f32 - }; - let conversation = column![ - scrollable(messages) - .id(chat_scroll_id()) - .height(Length::Fill), - container( - column![ - composer, - progress_bar(0.0..=1.0, context_fraction).height(3), - row![ - icon(ICON_PAPERCLIP, 19), - text(format!( - "{} / {} tokens ({:.0}%) • {}", - self.context_used, - self.context_limit, - context_fraction * 100.0, - self.tokens_per_second.map_or_else( - || "— tok/s".to_owned(), - |speed| format!("{speed:.1} tok/s") - ) - )) - .size(11) - .color(muted_text()), - Space::with_width(Length::Fill), - icon(ICON_MODEL, 16), - text( - ModelChoice::from_id(&self.preferences.selected_model) - .unwrap_or_default() - .to_string(), - ) - .size(12), - action, - ] - .align_y(Alignment::Center), - ] - .spacing(8), - ) - .padding(16) - .width(Length::Fill) - .style(overview_style), - ] - .height(Length::Fill) - .spacing(8); - container(conversation) - .max_width(860) - .center_x(Length::Fill) - .height(Length::Fill) - .into() - } else { - container( - column![ - text(if item.sessions.is_empty() { - "No sessions yet" - } else { - "Choose a session" - }) - .size(24), - text("Create a session above or select one from the sidebar.").size(14), - ] - .spacing(8) - .align_x(Alignment::Center), - ) - .center_x(Length::Fill) - .center_y(Length::Fill) - .into() - }; - - container(column![header, body].spacing(24)) - .width(Length::Fill) - .height(Length::Fill) - .padding(24) - .into() - } - - fn stats_dashboard(&self) -> Element<'_, Message> { - let stats = &self.metrics_snapshot; - let context_fraction = if stats.context_limit == 0 { - 0.0 - } else { - stats.context_used.min(stats.context_limit) as f32 / stats.context_limit as f32 - }; - let cache_fraction = if stats.last_prompt_tokens == 0 { - 0.0 - } else { - stats.last_cached_tokens as f32 / stats.last_prompt_tokens as f32 - }; - let cache_hit_fraction = if stats.kv_lookups == 0 { - 0.0 - } else { - stats.kv_hits as f32 / stats.kv_lookups as f32 - }; - let endpoint = if stats.server_listening { - format!("Listening · 127.0.0.1:{}", stats.server_port) - } else { - "Stopped".to_owned() - }; - let phase_color = match stats.phase { - crate::metrics::RuntimePhase::Generating => Color::from_rgb8(84, 170, 255), - crate::metrics::RuntimePhase::Prefilling | crate::metrics::RuntimePhase::Loading => { - Color::from_rgb8(240, 180, 70) - } - crate::metrics::RuntimePhase::Ready => Color::from_rgb8(72, 176, 112), - crate::metrics::RuntimePhase::Failed => Color::from_rgb8(220, 80, 86), - crate::metrics::RuntimePhase::Unloaded => muted_text(), - }; - let heading = container( - row![ - column![ - text("Runtime observability").size(24), - text(format!( - "{} · {} · uptime {}", - stats.model, - stats.source.label(), - format_duration(stats.uptime_seconds as f64) - )) - .size(12) - .color(muted_text()), - ] - .spacing(5), - Space::with_width(Length::Fill), - container(text(stats.phase.label()).size(12).color(phase_color)) - .padding([7, 11]) - .style(move |_| status_badge_style(phase_color)), - ] - .align_y(Alignment::Center), - ) - .padding(16) - .style(overview_style); - - let headline = column![ - row![ - metric_card( - "DECODE", - format!("{:.1} tok/s", stats.decode_tokens_per_second), - format!( - "{} completion tokens total", - format_count(stats.completion_tokens) - ), - ), - metric_card( - "PREFILL", - format!("{:.1} tok/s", stats.prefill_tokens_per_second), - format!("{} prompt tokens total", format_count(stats.prompt_tokens)), - ), - ] - .spacing(10), - row![ - metric_card( - "CONTEXT", - format!( - "{} / {}", - format_count(u64::from(stats.context_used)), - format_count(u64::from(stats.context_limit)) - ), - format!("{:.0}% occupied", context_fraction * 100.0), - ), - metric_card( - "WORK", - format!( - "{} active · {} queued", - stats.http_active, stats.queue_depth - ), - format!("{} runtime requests", format_count(stats.runtime_requests)), - ), - ] - .spacing(10), - ] - .spacing(10); - - let throughput = stats_panel( - "MODEL ACTIVITY · LAST 24 SECONDS", - column![ - mini_chart( - &self.metrics_history, - |point| point.decode_tokens_per_second, - Color::from_rgb8(84, 170, 255), - ), - row![ - text("Decode") - .size(12) - .color(Color::from_rgb8(84, 170, 255)), - Space::with_width(Length::Fill), - text(format!("{:.1} tok/s", stats.decode_tokens_per_second)) - .size(12) - .color(muted_text()), - ], - mini_chart( - &self.metrics_history, - |point| point.prefill_tokens_per_second, - Color::from_rgb8(157, 119, 255), - ), - row![ - text("Prefill") - .size(12) - .color(Color::from_rgb8(157, 119, 255)), - Space::with_width(Length::Fill), - text(format!("{:.1} tok/s", stats.prefill_tokens_per_second)) - .size(12) - .color(muted_text()), - ], - ] - .spacing(7) - .into(), - ); - let requests = stats_panel( - "SERVER REQUEST RATE · LAST 24 SECONDS", - column![ - mini_chart( - &self.metrics_history, - |point| point.http_requests_per_second, - Color::from_rgb8(72, 176, 112), - ), - row![ - text(format!("{} requests", format_count(stats.http_requests))).size(12), - Space::with_width(Length::Fill), - text(format!( - "{} errors · {} streaming", - stats.http_errors, stats.http_streaming_requests - )) - .size(12) - .color(muted_text()), - ], - ] - .spacing(7) - .into(), - ); - let latest = self.metrics_history.back().copied().unwrap_or_default(); - let kv_io = stats_panel( - "KV CHECKPOINT I/O · LAST 24 SECONDS", - column![ - mini_chart( - &self.metrics_history, - |point| point.kv_read_bytes_per_second, - Color::from_rgb8(67, 194, 203), - ), - row![ - text(if stats.kv_read_active { - "Disk read · active" - } else { - "Disk read" - }) - .size(12) - .color(Color::from_rgb8(67, 194, 203)), - Space::with_width(Length::Fill), - text(format_rate(latest.kv_read_bytes_per_second)) - .size(12) - .color(muted_text()), - ], - mini_chart( - &self.metrics_history, - |point| point.kv_write_bytes_per_second, - Color::from_rgb8(240, 180, 70), - ), - row![ - text(if stats.kv_write_active { - "Disk write · active" - } else { - "Disk write" - }) - .size(12) - .color(Color::from_rgb8(240, 180, 70)), - Space::with_width(Length::Fill), - text(format_rate(latest.kv_write_bytes_per_second)) - .size(12) - .color(muted_text()), - ], - ] - .spacing(7) - .into(), - ); - - let model = stats_panel( - "MODEL CORE", - column![ - metric_row("State", stats.phase.label()), - metric_row("Loaded model", stats.model), - metric_row("Mapped weights", format_bytes(stats.model_bytes)), - metric_row("Tensors", format_count(stats.tensor_count)), - metric_row("Vocabulary", format_count(stats.vocabulary_size)), - metric_row("Last load", format_milliseconds(stats.model_load_ms)), - metric_row( - "Lifecycle", - format!( - "{} loads · {} unloads", - stats.model_loads, stats.model_unloads - ), - ), - ] - .spacing(9) - .into(), - ); - let runtime = stats_panel( - "GENERATION", - column![ - metric_row("Last runtime", format_milliseconds(stats.last_runtime_ms)), - metric_row( - "Average runtime", - format_milliseconds(stats.average_runtime_ms) - ), - metric_row("Last prompt", format_count(stats.last_prompt_tokens)), - metric_row("Last reused", format_count(stats.last_cached_tokens)), - metric_row( - "Last completion", - format_count(stats.last_completion_tokens) - ), - metric_row("Cached tokens total", format_count(stats.cached_tokens)), - metric_row( - "Cache reuse", - format!("{:.0}% of last prompt", cache_fraction * 100.0), - ), - metric_row( - "Results", - format!( - "{} completed · {} failed", - stats.completed_requests, stats.failed_requests - ), - ), - ] - .spacing(9) - .into(), - ); - let cache = stats_panel( - "KV CACHE", - column![ - metric_row( - "Total", - format!( - "{} · {} files", - format_bytes(stats.kv_bytes), - stats.kv_files - ) - ), - metric_row( - "Local sessions", - format!( - "{} · {} files", - format_bytes(stats.local_kv_bytes), - stats.local_kv_files - ), - ), - metric_row( - "HTTP transient", - format!( - "{} · {} files", - format_bytes(stats.http_kv_bytes), - stats.http_kv_files - ), - ), - metric_row("Checkpoint writes", format_count(stats.checkpoint_writes)), - metric_row( - "Exact hits", - format!( - "{} · {} memory / {} disk", - stats.kv_hits, stats.kv_memory_hits, stats.kv_disk_hits - ), - ), - metric_row( - "Misses", - format!("{} · {} invalid", stats.kv_misses, stats.kv_invalid), - ), - metric_row("Lookups", format_count(stats.kv_lookups)), - metric_row( - "Exact hit rate", - format!("{:.1}%", cache_hit_fraction * 100.0) - ), - metric_row("Prefix hits", format_count(stats.kv_prefix_hits)), - metric_row( - "Reads", - format!( - "{} · {} · {} errors · last {}", - stats.kv_read_operations, - format_bytes(stats.kv_read_bytes), - stats.kv_read_errors, - format_milliseconds(stats.last_kv_read_ms), - ), - ), - metric_row( - "Writes", - format!( - "{} · {} · {} errors · last {}", - stats.kv_write_operations, - format_bytes(stats.kv_write_bytes), - stats.kv_write_errors, - format_milliseconds(stats.last_kv_write_ms), - ), - ), - progress_bar(0.0..=1.0, cache_fraction.min(1.0)).height(4), - ] - .spacing(9) - .into(), - ); - let server = stats_panel( - "LOCAL SERVER", - column![ - metric_row("Endpoint", endpoint), - metric_row("Active", format_count(u64::from(stats.http_active))), - metric_row("Completed", format_count(stats.http_completed)), - metric_row("Chat completions", format_count(stats.http_chat_requests)), - metric_row("Model queries", format_count(stats.http_model_requests)), - metric_row( - "Runtime sources", - format!( - "{} local · {} HTTP", - stats.local_requests, stats.endpoint_generations - ), - ), - metric_row("Received", format_bytes(stats.http_bytes_received)), - metric_row("Last latency", format_milliseconds(stats.last_http_ms)), - metric_row( - "Average latency", - format_milliseconds(stats.average_http_ms) - ), - ] - .spacing(9) - .into(), - ); - - scrollable( - container( - column![ - heading, - headline, - throughput, - kv_io, - requests, - row![model, runtime].spacing(10), - row![cache, server].spacing(10), - text("Counters are published by the runtime with relaxed atomics and sampled by the UI every 200 ms.") - .size(11) - .color(muted_text()), - ] - .spacing(12), - ) - .padding(24) - .max_width(960) - .center_x(Length::Fill), - ) - .height(Length::Fill) - .into() - } - - fn preferences_panel(&self) -> Element<'_, Message> { - let dspark_toggle: Option Message> = self - .preference_draft - .model - .supports_dspark() - .then_some(Message::PreferenceDsparkChanged); - let dspark = checkbox( - "Enable DSpark for this model", - self.preference_draft.dspark_enabled, - ) - .on_toggle_maybe(dspark_toggle); - let glm_mtp_toggle: Option Message> = (self.preference_draft.model - == ModelChoice::Glm52) - .then_some(Message::PreferenceGlmMtpChanged); - let glm_mtp_timing_toggle: Option Message> = (self.preference_draft.model - == ModelChoice::Glm52) - .then_some(Message::PreferenceGlmMtpTimingChanged); - let dspark_strict_toggle: Option Message> = self - .preference_draft - .model - .supports_dspark() - .then_some(Message::PreferenceDsparkStrictChanged); - let effective = self - .preference_draft - .generation() - .and_then(|generation| { - self.preference_draft.runtime().and_then(|runtime| { - crate::settings::effective_settings( - self.preference_draft.model, - &generation, - &runtime, - &models_path(), - ) - }) - }) - .ok(); - let engine = effective.as_ref().map(|settings| &settings.engine); - let turn = effective.as_ref().map(|settings| &settings.turn); - let mut power = text_input("100", &self.preference_draft.power_percent); - let mut prefill = text_input("Automatic", &self.preference_draft.prefill_chunk); - let mut ssd_full_layers = text_input("Automatic", &self.preference_draft.ssd_full_layers); - let mut steering_file = text_input( - "Direction-vector file path", - &self.preference_draft.directional_steering_file, - ); - let mut steering_ffn = - text_input("Automatic", &self.preference_draft.directional_steering_ffn); - let mut steering_attn = text_input("0", &self.preference_draft.directional_steering_attn); - let mut dspark_confidence = text_input( - "0.9 (DS4 default)", - &self.preference_draft.dspark_confidence_threshold, - ); - if self.preference_draft.model != ModelChoice::Glm52 { - power = power.on_input(Message::PreferencePowerChanged); - prefill = prefill.on_input(Message::PreferencePrefillChunkChanged); - steering_file = steering_file.on_input(Message::PreferenceSteeringFileChanged); - steering_ffn = steering_ffn.on_input(Message::PreferenceSteeringFfnChanged); - steering_attn = steering_attn.on_input(Message::PreferenceSteeringAttnChanged); - } else { - ssd_full_layers = ssd_full_layers.on_input(Message::PreferenceSsdFullLayersChanged); - } - if self.preference_draft.model.supports_dspark() { - dspark_confidence = - dspark_confidence.on_input(Message::PreferenceDsparkConfidenceChanged); - } - - let model_group = preference_group( - "MODEL & LIFECYCLE", - column![ - pick_list( - &MODEL_CHOICES[..], - Some(self.preference_draft.model), - Message::PreferenceModelChanged, - ) - .width(Length::Fill), - text(format!( - "Main: {}{}", - engine.map_or_else( - || "Invalid settings".to_owned(), - |engine| engine.artifacts.model.display().to_string(), - ), - engine - .and_then(|engine| engine.artifacts.mtp.as_ref()) - .map_or_else(String::new, |path| format!( - " • support: {}", - path.display() - )), - )) - .size(12), - row![ - text_input("10", &self.preference_draft.idle_timeout_minutes) - .on_input(Message::PreferenceTimeoutChanged) - .width(90) - .padding(9), - text("minutes before unloading the model").size(13), - ] - .spacing(10) - .align_y(Alignment::Center), - text("Enter a whole number from 1 to 1440.").size(12), - ] - .spacing(10), - ); - let endpoint_group = preference_group( - "LOCAL ENDPOINT", - column![ - preference_input_row( - "Port", - text_input("4000", &self.preference_draft.endpoint_port) - .on_input(Message::PreferenceEndpointPortChanged), - ), - text("Listens on 127.0.0.1. Saving a changed port restarts the local endpoint.") - .size(12), - ] - .spacing(10), - ); - let generation_group = preference_group( - "GENERATION", - column![ - preference_input_row( - "Context tokens", - text_input("32768", &self.preference_draft.context_tokens) - .on_input(Message::PreferenceContextChanged), - ), - preference_input_row( - "Maximum generated tokens", - text_input("50000", &self.preference_draft.max_generated_tokens) - .on_input(Message::PreferenceMaxTokensChanged), - ), - text("System prompt").size(13), - text_input( - "You are a helpful assistant", - &self.preference_draft.system_prompt, - ) - .on_input(Message::PreferenceSystemPromptChanged) - .padding(9), - Space::with_height(4), - text("SAMPLING & REASONING").size(11).color(muted_text()), - preference_input_row( - "Temperature", - text_input("DS4 default", &self.preference_draft.temperature) - .on_input(Message::PreferenceTemperatureChanged), - ), - preference_input_row( - "Top-p", - text_input("DS4 default", &self.preference_draft.top_p) - .on_input(Message::PreferenceTopPChanged), - ), - preference_input_row( - "Min-p", - text_input("DS4 default", &self.preference_draft.min_p) - .on_input(Message::PreferenceMinPChanged), - ), - preference_input_row( - "Seed", - text_input("Random", &self.preference_draft.seed) - .on_input(Message::PreferenceSeedChanged), - ), - row![ - text("Reasoning").size(13).width(Length::Fill), - pick_list( - &REASONING_MODES[..], - Some(self.preference_draft.reasoning_mode), - Message::PreferenceReasoningChanged, - ) - .width(240), - ] - .spacing(12) - .align_y(Alignment::Center), - text("Blank sampling values retain DS4's model-family defaults. Think Max needs at least 393216 context tokens.") - .size(12), - text(turn.map_or_else( - || "Effective settings will appear after valid values are entered.".to_owned(), - |settings| format!( - "Effective: {} context • {} max • temp {} • top-p {} • min-p {} • seed {} • {} • system prompt {}", - settings.context_tokens, - settings.max_generated_tokens, - settings.temperature, - settings.top_p, - settings.min_p, - settings.seed.map_or_else(|| "random".to_owned(), |seed| seed.to_string()), - settings.reasoning_mode, - if settings.system_prompt.is_empty() { "off" } else { "on" }, - ), - )) - .size(12), - ] - .spacing(10), - ); - let execution_group = preference_group( - "EXECUTION", - column![ - preference_input_row( - "CPU helper threads", - text_input("Automatic", &self.preference_draft.cpu_threads) - .on_input(Message::PreferenceCpuThreadsChanged), - ), - preference_input_row("GPU power percent", power), - preference_input_row("Prefill chunk", prefill), - checkbox("Prefer exact quality kernels", self.preference_draft.quality) - .on_toggle(Message::PreferenceQualityChanged), - checkbox("Warm mapped weights at load time", self.preference_draft.warm_weights) - .on_toggle(Message::PreferenceWarmWeightsChanged), - text(if self.preference_draft.model == ModelChoice::Glm52 { - "GLM 5.2 uses full GPU power and selects prefill chunks automatically." - } else { - "Blank numeric values preserve DS4's automatic engine behavior." - }) - .size(12), - text(engine.as_ref().map_or_else( - || "Effective execution settings will appear after valid values are entered." - .to_owned(), - |engine| { - let settings = engine.execution; - format!( - "Metal engine: threads {} • power {}% • prefill {} • quality {} • warm weights {}", - if settings.cpu_threads == 0 { "auto".to_owned() } else { settings.cpu_threads.to_string() }, - if settings.power_percent == 0 { 100 } else { settings.power_percent }, - if settings.prefill_chunk == 0 { "auto".to_owned() } else { settings.prefill_chunk.to_string() }, - if settings.quality { "on" } else { "off" }, - if settings.warm_weights { "on" } else { "off" }, - ) - }, - )) - .size(12), - ] - .spacing(10), - ); - let acceleration_group = preference_group( - "ACCELERATION & MEMORY", - column![ - text("SPECULATIVE DECODING").size(11).color(muted_text()), - preference_input_row( - "MTP draft tokens", - text_input("1", &self.preference_draft.mtp_draft_tokens) - .on_input(Message::PreferenceMtpDraftChanged), - ), - preference_input_row( - "MTP verifier margin", - text_input("3", &self.preference_draft.mtp_margin) - .on_input(Message::PreferenceMtpMarginChanged), - ), - checkbox("Enable integrated GLM MTP", self.preference_draft.glm_mtp) - .on_toggle_maybe(glm_mtp_toggle), - checkbox( - "Log GLM MTP timing counters", - self.preference_draft.glm_mtp_timing, - ) - .on_toggle_maybe(glm_mtp_timing_toggle), - dspark, - preference_input_row("DSpark confidence threshold", dspark_confidence), - checkbox( - "DSpark target-only decode", - self.preference_draft.dspark_strict, - ) - .on_toggle_maybe(dspark_strict_toggle), - text(if self.preference_draft.model.supports_dspark() { - "DSpark uses the managed support artifact; entering a threshold or enabling strict mode also enables DSpark." - } else if self.preference_draft.model == ModelChoice::Glm52 { - "GLM MTP is integrated; DSpark is unavailable for this model." - } else { - "No managed MTP support artifact is available for this model." - }) - .size(12), - text(engine.as_ref().map_or_else( - || "Effective speculative settings will appear after valid values are entered." - .to_owned(), - |engine| { - let settings = engine.speculative; - format!( - "Engine: MTP draft {} • margin {} • GLM MTP {} • timing {} • DSpark {} • confidence {}{} • target-only {}", - settings.mtp_draft_tokens, - settings.mtp_margin, - if settings.glm_mtp { "on" } else { "off" }, - if settings.glm_mtp_timing { "on" } else { "off" }, - if settings.dspark { "on" } else { "off" }, - settings.dspark_confidence_threshold, - if settings.dspark_confidence_threshold_set { " explicit" } else { " default" }, - if settings.dspark_strict { "on" } else { "off" }, - ) - }, - )) - .size(12), - Space::with_height(6), - text("SSD STREAMING").size(11).color(muted_text()), - checkbox("Enable SSD-backed model streaming", self.preference_draft.ssd_streaming) - .on_toggle(Message::PreferenceSsdChanged), - checkbox("Skip automatic expert preload", self.preference_draft.ssd_streaming_cold) - .on_toggle(Message::PreferenceSsdColdChanged), - preference_input_row( - "Expert cache count or GiB", - text_input("Automatic, 128, or 64GB", &self.preference_draft.ssd_cache) - .on_input(Message::PreferenceSsdCacheChanged), - ), - preference_input_row("Fully resident GLM layers", ssd_full_layers), - preference_input_row( - "Explicit expert preload count", - text_input("Automatic", &self.preference_draft.ssd_preload_experts) - .on_input(Message::PreferenceSsdPreloadChanged), - ), - text("A blank full-layer value is automatic; an explicit 0 disables fully resident GLM layers. SSD streaming and DSpark are mutually exclusive.") - .size(12), - text(engine.as_ref().map_or_else( - || "Effective SSD settings will appear after valid values are entered." - .to_owned(), - |engine| { - let settings = engine.ssd; - let cache = if settings.cache_bytes > 0 { - format!("{} GiB", settings.cache_bytes / GIB) - } else if settings.cache_experts > 0 { - format!("{} experts", settings.cache_experts) - } else { - "auto".to_owned() - }; - format!( - "Engine: streaming {} • cold {} • cache {} • full layers {}{} • preload {}", - if settings.enabled { "on" } else { "off" }, - if settings.cold { "on" } else { "off" }, - cache, - settings.full_layers, - if settings.full_layers_set { " explicit" } else { " auto" }, - if settings.preload_experts == 0 { "auto".to_owned() } else { settings.preload_experts.to_string() }, - ) - }, - )) - .size(12), - ] - .spacing(10), - ); - let steering_group = preference_group( - "STEERING & DIAGNOSTICS", - column![ - text("DIRECTIONAL STEERING").size(11).color(muted_text()), - text("Direction-vector file").size(13), - steering_file.padding(9), - preference_input_row("FFN scale", steering_ffn), - preference_input_row("Attention scale", steering_attn), - text(if self.preference_draft.model == ModelChoice::Glm52 { - "Directional steering is not supported for GLM 5.2." - } else { - "With a file and no explicit scale, DS4 defaults the FFN scale to 1. Scales accept -100 through 100." - }) - .size(12), - text(engine.as_ref().map_or_else( - || "Effective steering settings will appear after valid values are entered." - .to_owned(), - |engine| format!( - "Engine: file {} • FFN scale {} • attention scale {}", - if engine.steering.file.is_some() { "set" } else { "off" }, - engine.steering.ffn_scale, - engine.steering.attention_scale, - ), - )) - .size(12), - Space::with_height(6), - text("ADVANCED DIAGNOSTICS").size(11).color(muted_text()), - preference_input_row( - "Simulated used memory (GiB)", - text_input("Disabled", &self.preference_draft.simulated_used_memory_gib) - .on_input(Message::PreferenceSimulatedMemoryChanged), - ), - text("Routed expert profile output").size(13), - text_input("Output file path", &self.preference_draft.expert_profile_path) - .on_input(Message::PreferenceExpertProfileChanged) - .padding(9), - text(engine.as_ref().map_or_else( - || "Effective diagnostic settings will appear after valid values are entered." - .to_owned(), - |engine| format!( - "{} load: simulated memory {} • expert profile {}", - engine.model, - if engine.diagnostics.simulated_used_memory_bytes == 0 { - "off".to_owned() - } else { - format!("{} GiB", engine.diagnostics.simulated_used_memory_bytes / GIB) - }, - if engine.diagnostics.expert_profile_path.is_some() { "set" } else { "off" }, - ), - )) - .size(12), - ] - .spacing(10), - ); - let mut fields = column![ - model_group, - endpoint_group, - generation_group, - execution_group, - acceleration_group, - steering_group, - ] - .spacing(12); - - if let Some(error) = &self.preference_error { - fields = fields.push(text(error).style(iced::widget::text::danger)); - } - let header = row![ - icon(ICON_SETTINGS, 22), - text("Preferences").size(24), - Space::with_width(Length::Fill), - text("⌘,").size(12), - ] - .spacing(10) - .align_y(Alignment::Center); - let footer = row![ - action_button("Reset DS4 defaults").on_press(Message::ResetPreferences), - Space::with_width(Length::Fill), - action_button("Cancel").on_press(Message::DismissPanel), - action_button("Save").on_press(Message::SavePreferences), - ] - .spacing(8); - - let panel = container( - column![ - header, - scrollable(container(fields).padding(iced::Padding::ZERO.right(18))) - .height(Length::Fill), - footer - ] - .spacing(16), - ) - .padding(24) - .width(700) - .height(Length::Fill) - .max_height(660) - .style(overview_style); - opaque( - container(panel) - .padding(24) - .center_x(Length::Fill) - .center_y(Length::Fill) - .style(|_| { - container::Style::default().background(Color::from_rgba8(0, 0, 0, 0.68)) - }), - ) - } - - fn model_manager(&self) -> Element<'_, Message> { - let busy = matches!(self.model_download, ModelDownload::Active(_)); - let mut artifacts = column![]; - for (index, artifact) in model::managed_artifacts(&models_path()).iter().enumerate() { - if index > 0 { - artifacts = artifacts.push(horizontal_rule(1)); - } - artifacts = artifacts.push(model_artifact_row(artifact, busy)); - } - - let mut content = column![ - text("Model Manager").size(26), - text("Download, verify, or remove locally stored model files.") - .size(14) - .color(muted_text()), - ] - .spacing(8); - if !matches!(self.model_download, ModelDownload::Idle) { - content = content.push( - container(model_download_status(&self.model_download)) - .padding(16) - .style(overview_style), - ); - } - if let Some(error) = &self.error { - content = content.push(text(error).style(iced::widget::text::danger)); - } - content = content.push(Space::with_height(8)).push( - container(scrollable(artifacts).height(Length::Fill)) - .height(Length::Fill) - .style(overview_style), - ); - - let base: Element<'_, Message> = container(content) - .width(Length::Fill) - .height(Length::Fill) - .padding(28) - .into(); - let Some(artifact) = self.pending_model_delete else { - return base; - }; - let confirmation = container( - column![ - text("Delete model file?").size(22), - text(format!( - "Delete {artifact}, including any resumable partial download?" - )) - .size(13), - row![ - Space::with_width(Length::Fill), - action_button("Cancel").on_press(Message::CancelDeleteArtifact), - danger_button("Delete").on_press(Message::ConfirmDeleteArtifact), - ] - .spacing(8), - ] - .spacing(14), - ) - .padding(22) - .width(460) - .style(overview_style); - stack![ - base, - opaque( - container(confirmation) - .center_x(Length::Fill) - .center_y(Length::Fill) - .style(|_| container::Style::default() - .background(Color::from_rgba8(0, 0, 0, 0.68))) - ) - ] - .into() - } - fn project_dialog<'a>(&'a self, path: &'a Path) -> Element<'a, Message> { let dialog = container( column![ @@ -1358,62 +283,6 @@ impl App { } } -fn model_artifact_row(artifact: &ManagedArtifact, busy: bool) -> Element<'static, Message> { - let status = match artifact.state { - ManagedArtifactState::Missing => "Not downloaded", - ManagedArtifactState::Partial => "Partial download", - ManagedArtifactState::NeedsVerification => "Downloaded; verification required", - ManagedArtifactState::Ready => "Ready and verified", - }; - let download_label = match artifact.state { - ManagedArtifactState::Ready => "Downloaded", - _ if artifact.stored > 0 => "Resume", - _ => "Download", - }; - let download = if artifact.state == ManagedArtifactState::Ready || busy { - action_button(download_label) - } else { - action_button(download_label).on_press(Message::DownloadArtifact(artifact.id)) - }; - let validate = if artifact.can_validate() && !busy { - action_button("Validate").on_press(Message::ValidateArtifact(artifact.id)) - } else { - action_button("Validate") - }; - let delete = if artifact.stored > 0 && !busy { - danger_button("Delete").on_press(Message::DeleteArtifact(artifact.id)) - } else { - danger_button("Delete") - }; - - container( - row![ - icon(ICON_MODEL, 28), - column![ - text(artifact.id.to_string()).size(16), - text(status).size(13).color(muted_text()), - text(format!( - "{} on disk • {} expected", - format_bytes(artifact.stored), - format_bytes(artifact.expected), - )) - .size(12) - .color(muted_text()), - ] - .spacing(5), - Space::with_width(Length::Fill), - row![download, validate, delete] - .spacing(8) - .align_y(Alignment::Center), - ] - .spacing(16) - .align_y(Alignment::Center), - ) - .width(Length::Fill) - .padding([18, 20]) - .into() -} - fn preference_input_row<'a>( label: &'a str, input: iced::widget::TextInput<'a, Message>, @@ -1438,169 +307,6 @@ fn preference_group<'a>( .into() } -fn download_status_bar(download: &ActiveDownload) -> Element<'_, Message> { - let progress = &download.progress; - let percent = progress.fraction() * 100.0; - let measurement = if progress.verification.is_some() { - format!( - "{} verified of {} ({percent:.1}%)", - format_bytes(progress.completed()), - format_bytes(progress.active_total()), - ) - } else { - format!( - "{} of {} ({percent:.1}%)", - format_bytes(progress.completed()), - format_bytes(progress.active_total()), - ) - }; - let transfer = if download.bytes_per_second > 0.0 && progress.remaining() > 0 { - format!( - "{}/s • about {} remaining", - format_bytes(download.bytes_per_second as u64), - format_duration(progress.remaining() as f64 / download.bytes_per_second), - ) - } else { - "Calculating time remaining…".to_owned() - }; - let stop = if download.stopping { - danger_button("Stopping…") - } else { - danger_button("Stop").on_press(Message::StopModelDownload) - }; - - container( - row![ - text(phase_text(progress.phase)).size(12), - progress_bar(0.0..=1.0, progress.fraction()) - .width(180) - .height(7), - text(measurement).size(12), - text(transfer).size(12), - Space::with_width(Length::Fill), - stop, - ] - .spacing(12) - .align_y(Alignment::Center), - ) - .width(Length::Fill) - .padding([7, 14]) - .style(sidebar_style) - .into() -} - -fn model_download_status(download: &ModelDownload) -> Element<'_, Message> { - let (progress, heading, speed, failed) = match download { - ModelDownload::Idle => unreachable!("idle operations are not displayed"), - ModelDownload::Active(active) => ( - &active.progress, - if active.stopping { - format!("Stopping {}…", active.artifact) - } else { - phase_text(active.progress.phase) - }, - active.bytes_per_second, - false, - ), - ModelDownload::Complete(artifact, operation, progress) => ( - progress, - match operation { - ModelOperation::Download => format!("{artifact} is downloaded and verified."), - ModelOperation::Validate => format!("{artifact} passed validation."), - }, - 0.0, - false, - ), - ModelDownload::Failed(artifact, error, progress) => ( - progress, - format!("{artifact} operation failed: {error}"), - 0.0, - true, - ), - }; - - let percent = progress.fraction() * 100.0; - let measurement = if progress.verification.is_some() { - format!( - "{} verified of {} ({percent:.1}%) • {} remaining", - format_bytes(progress.completed()), - format_bytes(progress.active_total()), - format_bytes(progress.remaining()), - ) - } else { - format!( - "{} of {} ({percent:.1}%) • {} remaining", - format_bytes(progress.completed()), - format_bytes(progress.active_total()), - format_bytes(progress.remaining()), - ) - }; - let heading = if failed { - text(heading).size(12).style(iced::widget::text::danger) - } else { - text(heading).size(12) - }; - let mut heading_row = row![heading, Space::with_width(Length::Fill)] - .align_y(Alignment::Center) - .spacing(8); - if let ModelDownload::Active(active) = download { - heading_row = heading_row.push(if active.stopping { - danger_button("Stopping…") - } else { - danger_button("Stop").on_press(Message::StopModelDownload) - }); - } - let mut status = column![ - heading_row, - progress_bar(0.0..=1.0, progress.fraction()).height(8), - text(measurement).size(12), - ] - .spacing(6); - if speed > 0.0 && progress.remaining() > 0 { - status = status.push( - text(format!( - "{}/s • about {} remaining", - format_bytes(speed as u64), - format_duration(progress.remaining() as f64 / speed), - )) - .size(12), - ); - } - status.into() -} - -fn phase_text(phase: DownloadPhase) -> String { - match phase { - DownloadPhase::Pending(artifact) => format!("Ready to download {artifact}."), - DownloadPhase::Downloading(artifact) => format!("Downloading {artifact}…"), - DownloadPhase::Verifying(artifact) => format!("Verifying {artifact}…"), - DownloadPhase::Complete => "Download complete.".to_owned(), - } -} - -fn format_bytes(bytes: u64) -> String { - const GB: f64 = 1_000_000_000.0; - const MB: f64 = 1_000_000.0; - if bytes >= 1_000_000_000 { - format!("{:.1} GB", bytes as f64 / GB) - } else { - format!("{:.1} MB", bytes as f64 / MB) - } -} - -fn format_duration(seconds: f64) -> String { - let seconds = seconds.max(0.0).round() as u64; - let hours = seconds / 3600; - let minutes = seconds % 3600 / 60; - if hours > 0 { - format!("{hours}h {minutes}m") - } else if minutes > 0 { - format!("{minutes}m {}s", seconds % 60) - } else { - format!("{seconds}s") - } -} - pub(crate) fn app_theme() -> Theme { let palette = Palette { background: Color::from_rgb8(29, 29, 31), diff --git a/src/app/view/chat.rs b/src/app/view/chat.rs new file mode 100644 index 0000000..245ec40 --- /dev/null +++ b/src/app/view/chat.rs @@ -0,0 +1,210 @@ +use super::*; +use iced::widget::column; + +impl App { + pub(super) fn chat_detail(&self) -> Element<'_, Message> { + let Some(item) = self.selected_project() else { + let open_project_content = row![icon(ICON_FOLDER_PLUS, 17), text("Open project…"),] + .spacing(8) + .align_y(Alignment::Center); + let open_project = if self.database.is_some() && !self.choosing_folder { + action_button(open_project_content).on_press(Message::ChooseProjectFolder) + } else { + action_button(open_project_content) + }; + return container( + column![ + icon(ICON_SPARK, 36), + text("Start a local coding session").size(28), + text("Choose a project folder to create your first session.").size(14), + Space::with_height(10), + open_project, + ] + .spacing(10) + .align_x(Alignment::Center), + ) + .center_x(Length::Fill) + .center_y(Length::Fill) + .into(); + }; + + let project = &item.project; + let selected_title = self + .selected_session(item) + .map(|session| session.title.as_str()) + .unwrap_or(project.name.as_str()); + let header = row![ + icon(ICON_FOLDER, 19), + text(selected_title).size(18), + icon(ICON_MORE, 18), + Space::with_width(Length::Fill), + ] + .spacing(10) + .align_y(Alignment::Center); + + let body: Element<'_, Message> = if let Some(session) = self.selected_session(item) { + let mut messages = column![].spacing(12); + if self.conversation.is_empty() { + messages = messages.push( + column![ + text(&session.title).size(26), + text("Run DeepSeek locally with the Rust Metal engine.").size(14), + ] + .spacing(8), + ); + } else { + let markdown_style = markdown::Style::from_palette(app_theme().palette()); + for (index, message) in self.conversation.iter().enumerate() { + let label = if message.user { "You" } else { "DS4" }; + let active = self.generating && index + 1 == self.conversation.len(); + let mut body = column![text(label).size(11)].spacing(5); + if let Some(reasoning) = &message.reasoning { + let reasoning_label = + match (message.reasoning_open, message.reasoning_complete, active) { + (true, false, true) => "▾ Thinking", + (false, false, true) => "› Thinking", + (true, false, false) => "▾ Reasoning (stopped)", + (false, false, false) => "› Reasoning (stopped)", + (true, true, _) => "▾ Reasoning", + (false, true, _) => "› Reasoning", + }; + body = body.push( + button(text(reasoning_label).size(12)) + .padding(0) + .style(button::text) + .on_press(Message::ToggleReasoning(index)), + ); + if message.reasoning_open { + body = body.push( + text(if reasoning.is_empty() && active { + "Thinking…" + } else { + reasoning + }) + .size(13) + .color(muted_text()), + ); + } + } + if !message.content.is_empty() { + if message.user || message.markdown.is_empty() { + let content = if message.reasoning.is_some() { + message.content.trim_start() + } else { + &message.content + }; + body = body.push(text(content).size(14)); + } else { + body = body.push( + markdown::view( + &message.markdown, + markdown::Settings::with_text_size(14), + markdown_style, + ) + .map(Message::OpenLink), + ); + } + } else if active && message.reasoning.is_none() { + body = body.push(text("Loading model…").size(14)); + } + let user = message.user; + messages = messages.push( + container(body) + .padding(14) + .width(Length::Fill) + .style(move |theme| chat_message_style(theme, user)), + ); + } + } + let composer = text_input("Ask DS4Server anything…", &self.composer) + .on_input(Message::ComposerChanged) + .on_submit(Message::SubmitPrompt) + .padding(12) + .size(14); + let action = if self.generating { + action_button(text("Stop").size(12)).on_press(Message::StopGeneration) + } else if self.composer.trim().is_empty() { + action_button(icon(ICON_SEND, 18)).padding(8) + } else { + action_button(icon(ICON_SEND, 18)) + .padding(8) + .on_press(Message::SubmitPrompt) + }; + let context_fraction = if self.context_limit == 0 { + 0.0 + } else { + self.context_used.min(self.context_limit) as f32 / self.context_limit as f32 + }; + let conversation = column![ + scrollable(messages) + .id(chat_scroll_id()) + .height(Length::Fill), + container( + column![ + composer, + progress_bar(0.0..=1.0, context_fraction).height(3), + row![ + icon(ICON_PAPERCLIP, 19), + text(format!( + "{} / {} tokens ({:.0}%) • {}", + self.context_used, + self.context_limit, + context_fraction * 100.0, + self.tokens_per_second.map_or_else( + || "— tok/s".to_owned(), + |speed| format!("{speed:.1} tok/s") + ) + )) + .size(11) + .color(muted_text()), + Space::with_width(Length::Fill), + icon(ICON_MODEL, 16), + text( + ModelChoice::from_id(&self.preferences.selected_model) + .unwrap_or_default() + .to_string(), + ) + .size(12), + action, + ] + .align_y(Alignment::Center), + ] + .spacing(8), + ) + .padding(16) + .width(Length::Fill) + .style(overview_style), + ] + .height(Length::Fill) + .spacing(8); + container(conversation) + .max_width(860) + .center_x(Length::Fill) + .height(Length::Fill) + .into() + } else { + container( + column![ + text(if item.sessions.is_empty() { + "No sessions yet" + } else { + "Choose a session" + }) + .size(24), + text("Create a session above or select one from the sidebar.").size(14), + ] + .spacing(8) + .align_x(Alignment::Center), + ) + .center_x(Length::Fill) + .center_y(Length::Fill) + .into() + }; + + container(column![header, body].spacing(24)) + .width(Length::Fill) + .height(Length::Fill) + .padding(24) + .into() + } +} diff --git a/src/app/view/model_manager.rs b/src/app/view/model_manager.rs new file mode 100644 index 0000000..9226356 --- /dev/null +++ b/src/app/view/model_manager.rs @@ -0,0 +1,296 @@ +use super::*; +use iced::widget::column; + +impl App { + pub(super) fn model_manager(&self) -> Element<'_, Message> { + let busy = matches!(self.model_download, ModelDownload::Active(_)); + let mut artifacts = column![]; + for (index, artifact) in model::managed_artifacts(&models_path()).iter().enumerate() { + if index > 0 { + artifacts = artifacts.push(horizontal_rule(1)); + } + artifacts = artifacts.push(model_artifact_row(artifact, busy)); + } + + let mut content = column![ + text("Model Manager").size(26), + text("Download, verify, or remove locally stored model files.") + .size(14) + .color(muted_text()), + ] + .spacing(8); + if !matches!(self.model_download, ModelDownload::Idle) { + content = content.push( + container(model_download_status(&self.model_download)) + .padding(16) + .style(overview_style), + ); + } + if let Some(error) = &self.error { + content = content.push(text(error).style(iced::widget::text::danger)); + } + content = content.push(Space::with_height(8)).push( + container(scrollable(artifacts).height(Length::Fill)) + .height(Length::Fill) + .style(overview_style), + ); + + let base: Element<'_, Message> = container(content) + .width(Length::Fill) + .height(Length::Fill) + .padding(28) + .into(); + let Some(artifact) = self.pending_model_delete else { + return base; + }; + let confirmation = container( + column![ + text("Delete model file?").size(22), + text(format!( + "Delete {artifact}, including any resumable partial download?" + )) + .size(13), + row![ + Space::with_width(Length::Fill), + action_button("Cancel").on_press(Message::CancelDeleteArtifact), + danger_button("Delete").on_press(Message::ConfirmDeleteArtifact), + ] + .spacing(8), + ] + .spacing(14), + ) + .padding(22) + .width(460) + .style(overview_style); + stack![ + base, + opaque( + container(confirmation) + .center_x(Length::Fill) + .center_y(Length::Fill) + .style(|_| container::Style::default() + .background(Color::from_rgba8(0, 0, 0, 0.68))) + ) + ] + .into() + } +} + +fn model_artifact_row(artifact: &ManagedArtifact, busy: bool) -> Element<'static, Message> { + let status = match artifact.state { + ManagedArtifactState::Missing => "Not downloaded", + ManagedArtifactState::Partial => "Partial download", + ManagedArtifactState::NeedsVerification => "Downloaded; verification required", + ManagedArtifactState::Ready => "Ready and verified", + }; + let download_label = match artifact.state { + ManagedArtifactState::Ready => "Downloaded", + _ if artifact.stored > 0 => "Resume", + _ => "Download", + }; + let download = if artifact.state == ManagedArtifactState::Ready || busy { + action_button(download_label) + } else { + action_button(download_label).on_press(Message::DownloadArtifact(artifact.id)) + }; + let validate = if artifact.can_validate() && !busy { + action_button("Validate").on_press(Message::ValidateArtifact(artifact.id)) + } else { + action_button("Validate") + }; + let delete = if artifact.stored > 0 && !busy { + danger_button("Delete").on_press(Message::DeleteArtifact(artifact.id)) + } else { + danger_button("Delete") + }; + + container( + row![ + icon(ICON_MODEL, 28), + column![ + text(artifact.id.to_string()).size(16), + text(status).size(13).color(muted_text()), + text(format!( + "{} on disk • {} expected", + format_bytes(artifact.stored), + format_bytes(artifact.expected), + )) + .size(12) + .color(muted_text()), + ] + .spacing(5), + Space::with_width(Length::Fill), + row![download, validate, delete] + .spacing(8) + .align_y(Alignment::Center), + ] + .spacing(16) + .align_y(Alignment::Center), + ) + .width(Length::Fill) + .padding([18, 20]) + .into() +} + +pub(super) fn download_status_bar(download: &ActiveDownload) -> Element<'_, Message> { + let progress = &download.progress; + let percent = progress.fraction() * 100.0; + let measurement = if progress.verification.is_some() { + format!( + "{} verified of {} ({percent:.1}%)", + format_bytes(progress.completed()), + format_bytes(progress.active_total()), + ) + } else { + format!( + "{} of {} ({percent:.1}%)", + format_bytes(progress.completed()), + format_bytes(progress.active_total()), + ) + }; + let transfer = if download.bytes_per_second > 0.0 && progress.remaining() > 0 { + format!( + "{}/s • about {} remaining", + format_bytes(download.bytes_per_second as u64), + format_duration(progress.remaining() as f64 / download.bytes_per_second), + ) + } else { + "Calculating time remaining…".to_owned() + }; + let stop = if download.stopping { + danger_button("Stopping…") + } else { + danger_button("Stop").on_press(Message::StopModelDownload) + }; + + container( + row![ + text(phase_text(progress.phase)).size(12), + progress_bar(0.0..=1.0, progress.fraction()) + .width(180) + .height(7), + text(measurement).size(12), + text(transfer).size(12), + Space::with_width(Length::Fill), + stop, + ] + .spacing(12) + .align_y(Alignment::Center), + ) + .width(Length::Fill) + .padding([7, 14]) + .style(sidebar_style) + .into() +} + +fn model_download_status(download: &ModelDownload) -> Element<'_, Message> { + let (progress, heading, speed, failed) = match download { + ModelDownload::Idle => unreachable!("idle operations are not displayed"), + ModelDownload::Active(active) => ( + &active.progress, + if active.stopping { + format!("Stopping {}…", active.artifact) + } else { + phase_text(active.progress.phase) + }, + active.bytes_per_second, + false, + ), + ModelDownload::Complete(artifact, operation, progress) => ( + progress, + match operation { + ModelOperation::Download => format!("{artifact} is downloaded and verified."), + ModelOperation::Validate => format!("{artifact} passed validation."), + }, + 0.0, + false, + ), + ModelDownload::Failed(artifact, error, progress) => ( + progress, + format!("{artifact} operation failed: {error}"), + 0.0, + true, + ), + }; + + let percent = progress.fraction() * 100.0; + let measurement = if progress.verification.is_some() { + format!( + "{} verified of {} ({percent:.1}%) • {} remaining", + format_bytes(progress.completed()), + format_bytes(progress.active_total()), + format_bytes(progress.remaining()), + ) + } else { + format!( + "{} of {} ({percent:.1}%) • {} remaining", + format_bytes(progress.completed()), + format_bytes(progress.active_total()), + format_bytes(progress.remaining()), + ) + }; + let heading = if failed { + text(heading).size(12).style(iced::widget::text::danger) + } else { + text(heading).size(12) + }; + let mut heading_row = row![heading, Space::with_width(Length::Fill)] + .align_y(Alignment::Center) + .spacing(8); + if let ModelDownload::Active(active) = download { + heading_row = heading_row.push(if active.stopping { + danger_button("Stopping…") + } else { + danger_button("Stop").on_press(Message::StopModelDownload) + }); + } + let mut status = column![ + heading_row, + progress_bar(0.0..=1.0, progress.fraction()).height(8), + text(measurement).size(12), + ] + .spacing(6); + if speed > 0.0 && progress.remaining() > 0 { + status = status.push( + text(format!( + "{}/s • about {} remaining", + format_bytes(speed as u64), + format_duration(progress.remaining() as f64 / speed), + )) + .size(12), + ); + } + status.into() +} + +fn phase_text(phase: DownloadPhase) -> String { + match phase { + DownloadPhase::Pending(artifact) => format!("Ready to download {artifact}."), + DownloadPhase::Downloading(artifact) => format!("Downloading {artifact}…"), + DownloadPhase::Verifying(artifact) => format!("Verifying {artifact}…"), + DownloadPhase::Complete => "Download complete.".to_owned(), + } +} + +pub(super) fn format_bytes(bytes: u64) -> String { + const GB: f64 = 1_000_000_000.0; + const MB: f64 = 1_000_000.0; + if bytes >= 1_000_000_000 { + format!("{:.1} GB", bytes as f64 / GB) + } else { + format!("{:.1} MB", bytes as f64 / MB) + } +} + +pub(super) fn format_duration(seconds: f64) -> String { + let seconds = seconds.max(0.0).round() as u64; + let hours = seconds / 3600; + let minutes = seconds % 3600 / 60; + if hours > 0 { + format!("{hours}h {minutes}m") + } else if minutes > 0 { + format!("{minutes}m {}s", seconds % 60) + } else { + format!("{seconds}s") + } +} diff --git a/src/app/view/preferences.rs b/src/app/view/preferences.rs new file mode 100644 index 0000000..95fd4e4 --- /dev/null +++ b/src/app/view/preferences.rs @@ -0,0 +1,440 @@ +use super::*; +use iced::widget::column; + +impl App { + pub(super) fn preferences_panel(&self) -> Element<'_, Message> { + let dspark_toggle: Option Message> = self + .preference_draft + .model + .supports_dspark() + .then_some(Message::PreferenceDsparkChanged); + let dspark = checkbox( + "Enable DSpark for this model", + self.preference_draft.dspark_enabled, + ) + .on_toggle_maybe(dspark_toggle); + let glm_mtp_toggle: Option Message> = (self.preference_draft.model + == ModelChoice::Glm52) + .then_some(Message::PreferenceGlmMtpChanged); + let glm_mtp_timing_toggle: Option Message> = (self.preference_draft.model + == ModelChoice::Glm52) + .then_some(Message::PreferenceGlmMtpTimingChanged); + let dspark_strict_toggle: Option Message> = self + .preference_draft + .model + .supports_dspark() + .then_some(Message::PreferenceDsparkStrictChanged); + let effective = self + .preference_draft + .generation() + .and_then(|generation| { + self.preference_draft.runtime().and_then(|runtime| { + crate::settings::effective_settings( + self.preference_draft.model, + &generation, + &runtime, + &models_path(), + ) + }) + }) + .ok(); + let engine = effective.as_ref().map(|settings| &settings.engine); + let turn = effective.as_ref().map(|settings| &settings.turn); + let mut power = text_input("100", &self.preference_draft.power_percent); + let mut prefill = text_input("Automatic", &self.preference_draft.prefill_chunk); + let mut ssd_full_layers = text_input("Automatic", &self.preference_draft.ssd_full_layers); + let mut steering_file = text_input( + "Direction-vector file path", + &self.preference_draft.directional_steering_file, + ); + let mut steering_ffn = + text_input("Automatic", &self.preference_draft.directional_steering_ffn); + let mut steering_attn = text_input("0", &self.preference_draft.directional_steering_attn); + let mut dspark_confidence = text_input( + "0.9 (DS4 default)", + &self.preference_draft.dspark_confidence_threshold, + ); + if self.preference_draft.model != ModelChoice::Glm52 { + power = power.on_input(Message::PreferencePowerChanged); + prefill = prefill.on_input(Message::PreferencePrefillChunkChanged); + steering_file = steering_file.on_input(Message::PreferenceSteeringFileChanged); + steering_ffn = steering_ffn.on_input(Message::PreferenceSteeringFfnChanged); + steering_attn = steering_attn.on_input(Message::PreferenceSteeringAttnChanged); + } else { + ssd_full_layers = ssd_full_layers.on_input(Message::PreferenceSsdFullLayersChanged); + } + if self.preference_draft.model.supports_dspark() { + dspark_confidence = + dspark_confidence.on_input(Message::PreferenceDsparkConfidenceChanged); + } + + let model_group = preference_group( + "MODEL & LIFECYCLE", + column![ + pick_list( + &MODEL_CHOICES[..], + Some(self.preference_draft.model), + Message::PreferenceModelChanged, + ) + .width(Length::Fill), + text(format!( + "Main: {}{}", + engine.map_or_else( + || "Invalid settings".to_owned(), + |engine| engine.artifacts.model.display().to_string(), + ), + engine + .and_then(|engine| engine.artifacts.mtp.as_ref()) + .map_or_else(String::new, |path| format!( + " • support: {}", + path.display() + )), + )) + .size(12), + row![ + text_input("10", &self.preference_draft.idle_timeout_minutes) + .on_input(Message::PreferenceTimeoutChanged) + .width(90) + .padding(9), + text("minutes before unloading the model").size(13), + ] + .spacing(10) + .align_y(Alignment::Center), + text("Enter a whole number from 1 to 1440.").size(12), + ] + .spacing(10), + ); + let endpoint_group = preference_group( + "LOCAL ENDPOINT", + column![ + preference_input_row( + "Port", + text_input("4000", &self.preference_draft.endpoint_port) + .on_input(Message::PreferenceEndpointPortChanged), + ), + text("Listens on 127.0.0.1. Saving a changed port restarts the local endpoint.") + .size(12), + ] + .spacing(10), + ); + let generation_group = preference_group( + "GENERATION", + column![ + preference_input_row( + "Context tokens", + text_input("32768", &self.preference_draft.context_tokens) + .on_input(Message::PreferenceContextChanged), + ), + preference_input_row( + "Maximum generated tokens", + text_input("50000", &self.preference_draft.max_generated_tokens) + .on_input(Message::PreferenceMaxTokensChanged), + ), + text("System prompt").size(13), + text_input( + "You are a helpful assistant", + &self.preference_draft.system_prompt, + ) + .on_input(Message::PreferenceSystemPromptChanged) + .padding(9), + Space::with_height(4), + text("SAMPLING & REASONING").size(11).color(muted_text()), + preference_input_row( + "Temperature", + text_input("DS4 default", &self.preference_draft.temperature) + .on_input(Message::PreferenceTemperatureChanged), + ), + preference_input_row( + "Top-p", + text_input("DS4 default", &self.preference_draft.top_p) + .on_input(Message::PreferenceTopPChanged), + ), + preference_input_row( + "Min-p", + text_input("DS4 default", &self.preference_draft.min_p) + .on_input(Message::PreferenceMinPChanged), + ), + preference_input_row( + "Seed", + text_input("Random", &self.preference_draft.seed) + .on_input(Message::PreferenceSeedChanged), + ), + row![ + text("Reasoning").size(13).width(Length::Fill), + pick_list( + &REASONING_MODES[..], + Some(self.preference_draft.reasoning_mode), + Message::PreferenceReasoningChanged, + ) + .width(240), + ] + .spacing(12) + .align_y(Alignment::Center), + text("Blank sampling values retain DS4's model-family defaults. Think Max needs at least 393216 context tokens.") + .size(12), + text(turn.map_or_else( + || "Effective settings will appear after valid values are entered.".to_owned(), + |settings| format!( + "Effective: {} context • {} max • temp {} • top-p {} • min-p {} • seed {} • {} • system prompt {}", + settings.context_tokens, + settings.max_generated_tokens, + settings.temperature, + settings.top_p, + settings.min_p, + settings.seed.map_or_else(|| "random".to_owned(), |seed| seed.to_string()), + settings.reasoning_mode, + if settings.system_prompt.is_empty() { "off" } else { "on" }, + ), + )) + .size(12), + ] + .spacing(10), + ); + let execution_group = preference_group( + "EXECUTION", + column![ + preference_input_row( + "CPU helper threads", + text_input("Automatic", &self.preference_draft.cpu_threads) + .on_input(Message::PreferenceCpuThreadsChanged), + ), + preference_input_row("GPU power percent", power), + preference_input_row("Prefill chunk", prefill), + checkbox("Prefer exact quality kernels", self.preference_draft.quality) + .on_toggle(Message::PreferenceQualityChanged), + checkbox("Warm mapped weights at load time", self.preference_draft.warm_weights) + .on_toggle(Message::PreferenceWarmWeightsChanged), + text(if self.preference_draft.model == ModelChoice::Glm52 { + "GLM 5.2 uses full GPU power and selects prefill chunks automatically." + } else { + "Blank numeric values preserve DS4's automatic engine behavior." + }) + .size(12), + text(engine.as_ref().map_or_else( + || "Effective execution settings will appear after valid values are entered." + .to_owned(), + |engine| { + let settings = engine.execution; + format!( + "Metal engine: threads {} • power {}% • prefill {} • quality {} • warm weights {}", + if settings.cpu_threads == 0 { "auto".to_owned() } else { settings.cpu_threads.to_string() }, + if settings.power_percent == 0 { 100 } else { settings.power_percent }, + if settings.prefill_chunk == 0 { "auto".to_owned() } else { settings.prefill_chunk.to_string() }, + if settings.quality { "on" } else { "off" }, + if settings.warm_weights { "on" } else { "off" }, + ) + }, + )) + .size(12), + ] + .spacing(10), + ); + let acceleration_group = preference_group( + "ACCELERATION & MEMORY", + column![ + text("SPECULATIVE DECODING").size(11).color(muted_text()), + preference_input_row( + "MTP draft tokens", + text_input("1", &self.preference_draft.mtp_draft_tokens) + .on_input(Message::PreferenceMtpDraftChanged), + ), + preference_input_row( + "MTP verifier margin", + text_input("3", &self.preference_draft.mtp_margin) + .on_input(Message::PreferenceMtpMarginChanged), + ), + checkbox("Enable integrated GLM MTP", self.preference_draft.glm_mtp) + .on_toggle_maybe(glm_mtp_toggle), + checkbox( + "Log GLM MTP timing counters", + self.preference_draft.glm_mtp_timing, + ) + .on_toggle_maybe(glm_mtp_timing_toggle), + dspark, + preference_input_row("DSpark confidence threshold", dspark_confidence), + checkbox( + "DSpark target-only decode", + self.preference_draft.dspark_strict, + ) + .on_toggle_maybe(dspark_strict_toggle), + text(if self.preference_draft.model.supports_dspark() { + "DSpark uses the managed support artifact; entering a threshold or enabling strict mode also enables DSpark." + } else if self.preference_draft.model == ModelChoice::Glm52 { + "GLM MTP is integrated; DSpark is unavailable for this model." + } else { + "No managed MTP support artifact is available for this model." + }) + .size(12), + text(engine.as_ref().map_or_else( + || "Effective speculative settings will appear after valid values are entered." + .to_owned(), + |engine| { + let settings = engine.speculative; + format!( + "Engine: MTP draft {} • margin {} • GLM MTP {} • timing {} • DSpark {} • confidence {}{} • target-only {}", + settings.mtp_draft_tokens, + settings.mtp_margin, + if settings.glm_mtp { "on" } else { "off" }, + if settings.glm_mtp_timing { "on" } else { "off" }, + if settings.dspark { "on" } else { "off" }, + settings.dspark_confidence_threshold, + if settings.dspark_confidence_threshold_set { " explicit" } else { " default" }, + if settings.dspark_strict { "on" } else { "off" }, + ) + }, + )) + .size(12), + Space::with_height(6), + text("SSD STREAMING").size(11).color(muted_text()), + checkbox("Enable SSD-backed model streaming", self.preference_draft.ssd_streaming) + .on_toggle(Message::PreferenceSsdChanged), + checkbox("Skip automatic expert preload", self.preference_draft.ssd_streaming_cold) + .on_toggle(Message::PreferenceSsdColdChanged), + preference_input_row( + "Expert cache count or GiB", + text_input("Automatic, 128, or 64GB", &self.preference_draft.ssd_cache) + .on_input(Message::PreferenceSsdCacheChanged), + ), + preference_input_row("Fully resident GLM layers", ssd_full_layers), + preference_input_row( + "Explicit expert preload count", + text_input("Automatic", &self.preference_draft.ssd_preload_experts) + .on_input(Message::PreferenceSsdPreloadChanged), + ), + text("A blank full-layer value is automatic; an explicit 0 disables fully resident GLM layers. SSD streaming and DSpark are mutually exclusive.") + .size(12), + text(engine.as_ref().map_or_else( + || "Effective SSD settings will appear after valid values are entered." + .to_owned(), + |engine| { + let settings = engine.ssd; + let cache = if settings.cache_bytes > 0 { + format!("{} GiB", settings.cache_bytes / GIB) + } else if settings.cache_experts > 0 { + format!("{} experts", settings.cache_experts) + } else { + "auto".to_owned() + }; + format!( + "Engine: streaming {} • cold {} • cache {} • full layers {}{} • preload {}", + if settings.enabled { "on" } else { "off" }, + if settings.cold { "on" } else { "off" }, + cache, + settings.full_layers, + if settings.full_layers_set { " explicit" } else { " auto" }, + if settings.preload_experts == 0 { "auto".to_owned() } else { settings.preload_experts.to_string() }, + ) + }, + )) + .size(12), + ] + .spacing(10), + ); + let steering_group = preference_group( + "STEERING & DIAGNOSTICS", + column![ + text("DIRECTIONAL STEERING").size(11).color(muted_text()), + text("Direction-vector file").size(13), + steering_file.padding(9), + preference_input_row("FFN scale", steering_ffn), + preference_input_row("Attention scale", steering_attn), + text(if self.preference_draft.model == ModelChoice::Glm52 { + "Directional steering is not supported for GLM 5.2." + } else { + "With a file and no explicit scale, DS4 defaults the FFN scale to 1. Scales accept -100 through 100." + }) + .size(12), + text(engine.as_ref().map_or_else( + || "Effective steering settings will appear after valid values are entered." + .to_owned(), + |engine| format!( + "Engine: file {} • FFN scale {} • attention scale {}", + if engine.steering.file.is_some() { "set" } else { "off" }, + engine.steering.ffn_scale, + engine.steering.attention_scale, + ), + )) + .size(12), + Space::with_height(6), + text("ADVANCED DIAGNOSTICS").size(11).color(muted_text()), + preference_input_row( + "Simulated used memory (GiB)", + text_input("Disabled", &self.preference_draft.simulated_used_memory_gib) + .on_input(Message::PreferenceSimulatedMemoryChanged), + ), + text("Routed expert profile output").size(13), + text_input("Output file path", &self.preference_draft.expert_profile_path) + .on_input(Message::PreferenceExpertProfileChanged) + .padding(9), + text(engine.as_ref().map_or_else( + || "Effective diagnostic settings will appear after valid values are entered." + .to_owned(), + |engine| format!( + "{} load: simulated memory {} • expert profile {}", + engine.model, + if engine.diagnostics.simulated_used_memory_bytes == 0 { + "off".to_owned() + } else { + format!("{} GiB", engine.diagnostics.simulated_used_memory_bytes / GIB) + }, + if engine.diagnostics.expert_profile_path.is_some() { "set" } else { "off" }, + ), + )) + .size(12), + ] + .spacing(10), + ); + let mut fields = column![ + model_group, + endpoint_group, + generation_group, + execution_group, + acceleration_group, + steering_group, + ] + .spacing(12); + + if let Some(error) = &self.preference_error { + fields = fields.push(text(error).style(iced::widget::text::danger)); + } + let header = row![ + icon(ICON_SETTINGS, 22), + text("Preferences").size(24), + Space::with_width(Length::Fill), + text("⌘,").size(12), + ] + .spacing(10) + .align_y(Alignment::Center); + let footer = row![ + action_button("Reset DS4 defaults").on_press(Message::ResetPreferences), + Space::with_width(Length::Fill), + action_button("Cancel").on_press(Message::DismissPanel), + action_button("Save").on_press(Message::SavePreferences), + ] + .spacing(8); + + let panel = container( + column![ + header, + scrollable(container(fields).padding(iced::Padding::ZERO.right(18))) + .height(Length::Fill), + footer + ] + .spacing(16), + ) + .padding(24) + .width(700) + .height(Length::Fill) + .max_height(660) + .style(overview_style); + opaque( + container(panel) + .padding(24) + .center_x(Length::Fill) + .center_y(Length::Fill) + .style(|_| { + container::Style::default().background(Color::from_rgba8(0, 0, 0, 0.68)) + }), + ) + } +} diff --git a/src/app/view/stats.rs b/src/app/view/stats.rs new file mode 100644 index 0000000..150e5ce --- /dev/null +++ b/src/app/view/stats.rs @@ -0,0 +1,371 @@ +use super::*; +use iced::widget::column; + +impl App { + pub(super) fn stats_dashboard(&self) -> Element<'_, Message> { + let stats = &self.metrics_snapshot; + let context_fraction = if stats.context_limit == 0 { + 0.0 + } else { + stats.context_used.min(stats.context_limit) as f32 / stats.context_limit as f32 + }; + let cache_fraction = if stats.last_prompt_tokens == 0 { + 0.0 + } else { + stats.last_cached_tokens as f32 / stats.last_prompt_tokens as f32 + }; + let cache_hit_fraction = if stats.kv_lookups == 0 { + 0.0 + } else { + stats.kv_hits as f32 / stats.kv_lookups as f32 + }; + let endpoint = if stats.server_listening { + format!("Listening · 127.0.0.1:{}", stats.server_port) + } else { + "Stopped".to_owned() + }; + let phase_color = match stats.phase { + crate::metrics::RuntimePhase::Generating => Color::from_rgb8(84, 170, 255), + crate::metrics::RuntimePhase::Prefilling | crate::metrics::RuntimePhase::Loading => { + Color::from_rgb8(240, 180, 70) + } + crate::metrics::RuntimePhase::Ready => Color::from_rgb8(72, 176, 112), + crate::metrics::RuntimePhase::Failed => Color::from_rgb8(220, 80, 86), + crate::metrics::RuntimePhase::Unloaded => muted_text(), + }; + let heading = container( + row![ + column![ + text("Runtime observability").size(24), + text(format!( + "{} · {} · uptime {}", + stats.model, + stats.source.label(), + format_duration(stats.uptime_seconds as f64) + )) + .size(12) + .color(muted_text()), + ] + .spacing(5), + Space::with_width(Length::Fill), + container(text(stats.phase.label()).size(12).color(phase_color)) + .padding([7, 11]) + .style(move |_| status_badge_style(phase_color)), + ] + .align_y(Alignment::Center), + ) + .padding(16) + .style(overview_style); + + let headline = column![ + row![ + metric_card( + "DECODE", + format!("{:.1} tok/s", stats.decode_tokens_per_second), + format!( + "{} completion tokens total", + format_count(stats.completion_tokens) + ), + ), + metric_card( + "PREFILL", + format!("{:.1} tok/s", stats.prefill_tokens_per_second), + format!("{} prompt tokens total", format_count(stats.prompt_tokens)), + ), + ] + .spacing(10), + row![ + metric_card( + "CONTEXT", + format!( + "{} / {}", + format_count(u64::from(stats.context_used)), + format_count(u64::from(stats.context_limit)) + ), + format!("{:.0}% occupied", context_fraction * 100.0), + ), + metric_card( + "WORK", + format!( + "{} active · {} queued", + stats.http_active, stats.queue_depth + ), + format!("{} runtime requests", format_count(stats.runtime_requests)), + ), + ] + .spacing(10), + ] + .spacing(10); + + let throughput = stats_panel( + "MODEL ACTIVITY · LAST 24 SECONDS", + column![ + mini_chart( + &self.metrics_history, + |point| point.decode_tokens_per_second, + Color::from_rgb8(84, 170, 255), + ), + row![ + text("Decode") + .size(12) + .color(Color::from_rgb8(84, 170, 255)), + Space::with_width(Length::Fill), + text(format!("{:.1} tok/s", stats.decode_tokens_per_second)) + .size(12) + .color(muted_text()), + ], + mini_chart( + &self.metrics_history, + |point| point.prefill_tokens_per_second, + Color::from_rgb8(157, 119, 255), + ), + row![ + text("Prefill") + .size(12) + .color(Color::from_rgb8(157, 119, 255)), + Space::with_width(Length::Fill), + text(format!("{:.1} tok/s", stats.prefill_tokens_per_second)) + .size(12) + .color(muted_text()), + ], + ] + .spacing(7) + .into(), + ); + let requests = stats_panel( + "SERVER REQUEST RATE · LAST 24 SECONDS", + column![ + mini_chart( + &self.metrics_history, + |point| point.http_requests_per_second, + Color::from_rgb8(72, 176, 112), + ), + row![ + text(format!("{} requests", format_count(stats.http_requests))).size(12), + Space::with_width(Length::Fill), + text(format!( + "{} errors · {} streaming", + stats.http_errors, stats.http_streaming_requests + )) + .size(12) + .color(muted_text()), + ], + ] + .spacing(7) + .into(), + ); + let latest = self.metrics_history.back().copied().unwrap_or_default(); + let kv_io = stats_panel( + "KV CHECKPOINT I/O · LAST 24 SECONDS", + column![ + mini_chart( + &self.metrics_history, + |point| point.kv_read_bytes_per_second, + Color::from_rgb8(67, 194, 203), + ), + row![ + text(if stats.kv_read_active { + "Disk read · active" + } else { + "Disk read" + }) + .size(12) + .color(Color::from_rgb8(67, 194, 203)), + Space::with_width(Length::Fill), + text(format_rate(latest.kv_read_bytes_per_second)) + .size(12) + .color(muted_text()), + ], + mini_chart( + &self.metrics_history, + |point| point.kv_write_bytes_per_second, + Color::from_rgb8(240, 180, 70), + ), + row![ + text(if stats.kv_write_active { + "Disk write · active" + } else { + "Disk write" + }) + .size(12) + .color(Color::from_rgb8(240, 180, 70)), + Space::with_width(Length::Fill), + text(format_rate(latest.kv_write_bytes_per_second)) + .size(12) + .color(muted_text()), + ], + ] + .spacing(7) + .into(), + ); + + let model = stats_panel( + "MODEL CORE", + column![ + metric_row("State", stats.phase.label()), + metric_row("Loaded model", stats.model), + metric_row("Mapped weights", format_bytes(stats.model_bytes)), + metric_row("Tensors", format_count(stats.tensor_count)), + metric_row("Vocabulary", format_count(stats.vocabulary_size)), + metric_row("Last load", format_milliseconds(stats.model_load_ms)), + metric_row( + "Lifecycle", + format!( + "{} loads · {} unloads", + stats.model_loads, stats.model_unloads + ), + ), + ] + .spacing(9) + .into(), + ); + let runtime = stats_panel( + "GENERATION", + column![ + metric_row("Last runtime", format_milliseconds(stats.last_runtime_ms)), + metric_row( + "Average runtime", + format_milliseconds(stats.average_runtime_ms) + ), + metric_row("Last prompt", format_count(stats.last_prompt_tokens)), + metric_row("Last reused", format_count(stats.last_cached_tokens)), + metric_row( + "Last completion", + format_count(stats.last_completion_tokens) + ), + metric_row("Cached tokens total", format_count(stats.cached_tokens)), + metric_row( + "Cache reuse", + format!("{:.0}% of last prompt", cache_fraction * 100.0), + ), + metric_row( + "Results", + format!( + "{} completed · {} failed", + stats.completed_requests, stats.failed_requests + ), + ), + ] + .spacing(9) + .into(), + ); + let cache = stats_panel( + "KV CACHE", + column![ + metric_row( + "Total", + format!( + "{} · {} files", + format_bytes(stats.kv_bytes), + stats.kv_files + ) + ), + metric_row( + "Local sessions", + format!( + "{} · {} files", + format_bytes(stats.local_kv_bytes), + stats.local_kv_files + ), + ), + metric_row( + "HTTP transient", + format!( + "{} · {} files", + format_bytes(stats.http_kv_bytes), + stats.http_kv_files + ), + ), + metric_row("Checkpoint writes", format_count(stats.checkpoint_writes)), + metric_row( + "Exact hits", + format!( + "{} · {} memory / {} disk", + stats.kv_hits, stats.kv_memory_hits, stats.kv_disk_hits + ), + ), + metric_row( + "Misses", + format!("{} · {} invalid", stats.kv_misses, stats.kv_invalid), + ), + metric_row("Lookups", format_count(stats.kv_lookups)), + metric_row( + "Exact hit rate", + format!("{:.1}%", cache_hit_fraction * 100.0) + ), + metric_row("Prefix hits", format_count(stats.kv_prefix_hits)), + metric_row( + "Reads", + format!( + "{} · {} · {} errors · last {}", + stats.kv_read_operations, + format_bytes(stats.kv_read_bytes), + stats.kv_read_errors, + format_milliseconds(stats.last_kv_read_ms), + ), + ), + metric_row( + "Writes", + format!( + "{} · {} · {} errors · last {}", + stats.kv_write_operations, + format_bytes(stats.kv_write_bytes), + stats.kv_write_errors, + format_milliseconds(stats.last_kv_write_ms), + ), + ), + progress_bar(0.0..=1.0, cache_fraction.min(1.0)).height(4), + ] + .spacing(9) + .into(), + ); + let server = stats_panel( + "LOCAL SERVER", + column![ + metric_row("Endpoint", endpoint), + metric_row("Active", format_count(u64::from(stats.http_active))), + metric_row("Completed", format_count(stats.http_completed)), + metric_row("Chat completions", format_count(stats.http_chat_requests)), + metric_row("Model queries", format_count(stats.http_model_requests)), + metric_row( + "Runtime sources", + format!( + "{} local · {} HTTP", + stats.local_requests, stats.endpoint_generations + ), + ), + metric_row("Received", format_bytes(stats.http_bytes_received)), + metric_row("Last latency", format_milliseconds(stats.last_http_ms)), + metric_row( + "Average latency", + format_milliseconds(stats.average_http_ms) + ), + ] + .spacing(9) + .into(), + ); + + scrollable( + container( + column![ + heading, + headline, + throughput, + kv_io, + requests, + row![model, runtime].spacing(10), + row![cache, server].spacing(10), + text("Counters are published by the runtime with relaxed atomics and sampled by the UI every 200 ms.") + .size(11) + .color(muted_text()), + ] + .spacing(12), + ) + .padding(24) + .max_width(960) + .center_x(Length::Fill), + ) + .height(Length::Fill) + .into() + } +} diff --git a/src/engine.rs b/src/engine.rs index 7d036bd..6739b9b 100644 --- a/src/engine.rs +++ b/src/engine.rs @@ -2,6 +2,7 @@ mod gguf; #[cfg(target_os = "macos")] mod metal; mod tokenizer; +mod validation; #[cfg(target_os = "macos")] use crate::metrics::{KvLookup, Metrics}; @@ -17,6 +18,9 @@ use std::sync::atomic::{AtomicBool, Ordering}; #[cfg(target_os = "macos")] use std::time::Instant; use tokenizer::Tokenizer; +use validation::{validate_dspark, validate_main}; + +pub(crate) use validation::validate_model_artifact; #[cfg(target_os = "macos")] pub(crate) use metal::configure_sources as configure_metal_sources; @@ -1086,996 +1090,3 @@ mod sampling_tests { std::fs::remove_file(checkpoint).unwrap(); } } - -pub(crate) fn validate_model_artifact( - path: &Path, - expected: ModelChoice, - support: bool, -) -> Result<(), String> { - if support { - let model = Gguf::open(path)?; - validate_dspark(&model, &FLASH) - } else { - let model = Model::open_main(path, expected)?; - let summary = model.summary(); - if summary.tensor_count == 0 - || model - .render_prompt("", "", ReasoningMode::Direct) - .is_empty() - { - return Err("model intake produced an empty tensor directory or prompt".into()); - } - let _ = model.tokenize(""); - let eos = model.eos_token(); - let _ = model.token_bytes(eos); - if !model.is_stop_token(eos) { - return Err("tokenizer EOS marker is not a stop token".into()); - } - model.tensor_data("token_embd.weight")?; - Ok(()) - } -} - -fn validate_main(model: &Gguf, expected: ModelChoice) -> Result { - let family = if model.bytes("general.architecture").ok() == Some(b"glm-dsa") { - ModelFamily::Glm - } else { - ModelFamily::DeepSeek - }; - let shape = match family { - ModelFamily::Glm => GLM, - ModelFamily::DeepSeek => match model.u32("deepseek4.block_count")? { - 43 => FLASH, - 61 => PRO, - layers => return Err(format!("unsupported DeepSeek layer count: {layers}")), - }, - }; - if shape.model != expected { - return Err(format!( - "{} contains {}, but the selected model is {expected}", - model.path().display(), - shape.model - )); - } - validate_metadata(model, &shape)?; - validate_tensors(model, &shape)?; - Ok(shape) -} - -fn validate_metadata(model: &Gguf, shape: &Shape) -> Result<(), String> { - let prefix = if shape.family == ModelFamily::Glm { - "glm-dsa" - } else { - "deepseek4" - }; - for (key, expected) in [ - ("block_count", u64::from(shape.layers)), - ("embedding_length", shape.embd), - ("vocab_size", shape.vocab), - ("attention.head_count", shape.heads), - ("attention.head_count_kv", shape.head_kv), - ("attention.key_length", shape.head_dim), - ("attention.value_length", shape.value_dim), - ("rope.dimension_count", shape.rot), - ("attention.q_lora_rank", shape.lora_q), - ("expert_count", shape.experts), - ("expert_used_count", shape.experts_used), - ("expert_feed_forward_length", shape.ff_expert), - ("expert_shared_count", shape.expert_shared), - ("attention.indexer.head_count", shape.indexer_heads), - ("attention.indexer.key_length", shape.indexer_head_dim), - ("attention.indexer.top_k", shape.indexer_top_k), - ] { - expect_u64(model, &format!("{prefix}.{key}"), expected)?; - } - expect_float( - model, - &format!("{prefix}.attention.layer_norm_rms_epsilon"), - shape.rms_epsilon, - )?; - expect_float( - model, - &format!("{prefix}.expert_weights_scale"), - shape.expert_weight_scale, - )?; - if !model.boolean(&format!("{prefix}.expert_weights_norm"))? { - return Err(format!("{prefix}.expert_weights_norm must be true")); - } - expect_float(model, &format!("{prefix}.rope.freq_base"), shape.rope_base)?; - - if shape.family == ModelFamily::Glm { - for (key, expected) in [ - ("context_length", shape.original_context), - ("feed_forward_length", shape.ff_dense), - ("attention.kv_lora_rank", shape.kv_lora), - ("attention.key_length_mla", shape.key_mla), - ("attention.value_length_mla", shape.value_mla), - ("expert_group_count", 1), - ("expert_group_used_count", 1), - ("expert_gating_func", 2), - ("leading_dense_block_count", u64::from(shape.leading_dense)), - ("nextn_predict_layers", u64::from(shape.nextn)), - ] { - expect_u64(model, &format!("{prefix}.{key}"), expected)?; - } - return Ok(()); - } - - for (key, expected) in [ - ("attention.output_group_count", shape.out_groups), - ("attention.output_lora_rank", shape.lora_o), - ("hash_layer_count", u64::from(shape.hash_layers)), - ("attention.sliding_window", shape.sliding_window), - ("hyper_connection.count", shape.hc), - ("hyper_connection.sinkhorn_iterations", shape.hc_sinkhorn), - ] { - expect_u64(model, &format!("{prefix}.{key}"), expected)?; - } - for key in ["expert_group_count", "expert_group_used_count"] { - if let Some(Value::U32(value)) = model.metadata.get(&format!("{prefix}.{key}")) - && *value != 0 - { - return Err(format!("{prefix}.{key} must be zero")); - } - } - for (key, expected) in [ - ("hyper_connection.epsilon", shape.hc_epsilon), - ( - "attention.compress_rope_freq_base", - shape.compress_rope_base, - ), - ] { - expect_float(model, &format!("{prefix}.{key}"), expected)?; - } - for (key, expected) in [ - ("rope.scaling.factor", shape.rope_scale), - ("rope.scaling.yarn_beta_fast", shape.rope_beta_fast), - ("rope.scaling.yarn_beta_slow", shape.rope_beta_slow), - ] { - if model.metadata.contains_key(&format!("{prefix}.{key}")) { - expect_float(model, &format!("{prefix}.{key}"), expected)?; - } - } - if model - .metadata - .contains_key("deepseek4.rope.scaling.original_context_length") - { - expect_u64( - model, - "deepseek4.rope.scaling.original_context_length", - shape.original_context, - )?; - } - let ratios = model.u32s("deepseek4.attention.compress_ratios")?; - let clamps = model.f32s("deepseek4.swiglu_clamp_exp")?; - if ratios.len() < shape.layers as usize || clamps.len() < shape.layers as usize { - return Err("DeepSeek per-layer metadata is shorter than the layer count".into()); - } - for layer in 0..shape.layers as usize { - let expected = compression_ratio(shape, layer as u32); - if ratios[layer] != expected { - return Err(format!( - "layer {layer} compression ratio is {}, expected {expected}", - ratios[layer] - )); - } - if !float_eq(clamps[layer], shape.swiglu_clamp) { - return Err(format!("layer {layer} has an invalid SwiGLU clamp")); - } - } - Ok(()) -} - -fn validate_tensors(model: &Gguf, shape: &Shape) -> Result<(), String> { - match shape.family { - ModelFamily::DeepSeek => validate_deepseek_tensors(model, shape), - ModelFamily::Glm => validate_glm_tensors(model, shape), - } -} - -fn validate_deepseek_tensors(model: &Gguf, shape: &Shape) -> Result<(), String> { - let hc_dim = shape.embd * shape.hc; - let hc_mix = 2 * shape.hc + shape.hc * shape.hc; - let q_dim = shape.heads * shape.head_dim; - let output_low = shape.out_groups * shape.lora_o; - expect( - model, - "token_embd.weight", - &[F16], - &[shape.embd, shape.vocab], - )?; - expect(model, "output_hc_base.weight", &[F32], &[shape.hc])?; - expect(model, "output_hc_fn.weight", &[F16], &[hc_dim, shape.hc])?; - expect(model, "output_hc_scale.weight", &[F32], &[1])?; - expect(model, "output_norm.weight", &[F32], &[shape.embd])?; - expect(model, "output.weight", DENSE, &[shape.embd, shape.vocab])?; - - for layer in 0..shape.layers { - let name = |suffix: &str| format!("blk.{layer}.{suffix}"); - expect(model, &name("hc_attn_fn.weight"), &[F16], &[hc_dim, hc_mix])?; - expect(model, &name("hc_attn_scale.weight"), &[F32], &[3])?; - expect(model, &name("hc_attn_base.weight"), &[F32], &[hc_mix])?; - expect(model, &name("attn_norm.weight"), &[F32], &[shape.embd])?; - expect( - model, - &name("attn_q_a.weight"), - DENSE, - &[shape.embd, shape.lora_q], - )?; - expect( - model, - &name("attn_q_a_norm.weight"), - &[F32], - &[shape.lora_q], - )?; - expect( - model, - &name("attn_q_b.weight"), - DENSE, - &[shape.lora_q, q_dim], - )?; - expect( - model, - &name("attn_kv.weight"), - DENSE, - &[shape.embd, shape.head_dim], - )?; - expect( - model, - &name("attn_kv_a_norm.weight"), - &[F32], - &[shape.head_dim], - )?; - expect(model, &name("attn_sinks.weight"), &[F32], &[shape.heads])?; - expect( - model, - &name("attn_output_a.weight"), - DENSE, - &[ - shape.head_dim * (shape.heads / shape.out_groups), - output_low, - ], - )?; - expect( - model, - &name("attn_output_b.weight"), - DENSE, - &[output_low, shape.embd], - )?; - - let ratio = compression_ratio(shape, layer); - if ratio != 0 { - let compression_width = if ratio == 4 { 2 } else { 1 } * shape.head_dim; - expect( - model, - &name("attn_compressor_ape.weight"), - &[F16], - &[compression_width, u64::from(ratio)], - )?; - expect( - model, - &name("attn_compressor_kv.weight"), - &[F16], - &[shape.embd, compression_width], - )?; - expect( - model, - &name("attn_compressor_gate.weight"), - &[F16], - &[shape.embd, compression_width], - )?; - expect( - model, - &name("attn_compressor_norm.weight"), - &[F32], - &[shape.head_dim], - )?; - } - if ratio == 4 { - let index_q = shape.indexer_heads * shape.indexer_head_dim; - let index_width = 2 * shape.indexer_head_dim; - expect( - model, - &name("indexer.attn_q_b.weight"), - &[F16, Q8_0], - &[shape.lora_q, index_q], - )?; - expect( - model, - &name("indexer.proj.weight"), - &[F16], - &[shape.embd, shape.indexer_heads], - )?; - expect( - model, - &name("indexer_compressor_ape.weight"), - &[F16], - &[index_width, 4], - )?; - expect( - model, - &name("indexer_compressor_kv.weight"), - &[F16], - &[shape.embd, index_width], - )?; - expect( - model, - &name("indexer_compressor_gate.weight"), - &[F16], - &[shape.embd, index_width], - )?; - expect( - model, - &name("indexer_compressor_norm.weight"), - &[F32], - &[shape.indexer_head_dim], - )?; - } - expect(model, &name("hc_ffn_fn.weight"), &[F16], &[hc_dim, hc_mix])?; - expect(model, &name("hc_ffn_scale.weight"), &[F32], &[3])?; - expect(model, &name("hc_ffn_base.weight"), &[F32], &[hc_mix])?; - expect(model, &name("ffn_norm.weight"), &[F32], &[shape.embd])?; - expect( - model, - &name("ffn_gate_inp.weight"), - &[F16], - &[shape.embd, shape.experts], - )?; - expect_optional(model, &name("exp_probs_b.bias"), &[F32], &[shape.experts])?; - expect( - model, - &name("ffn_gate_exps.weight"), - ROUTED, - &[shape.embd, shape.ff_expert, shape.experts], - )?; - expect( - model, - &name("ffn_up_exps.weight"), - ROUTED, - &[shape.embd, shape.ff_expert, shape.experts], - )?; - expect( - model, - &name("ffn_down_exps.weight"), - ROUTED, - &[shape.ff_expert, shape.embd, shape.experts], - )?; - same_type( - model, - &name("ffn_gate_exps.weight"), - &name("ffn_up_exps.weight"), - )?; - expect( - model, - &name("ffn_gate_shexp.weight"), - DENSE, - &[shape.embd, shape.ff_expert], - )?; - expect( - model, - &name("ffn_up_shexp.weight"), - DENSE, - &[shape.embd, shape.ff_expert], - )?; - expect( - model, - &name("ffn_down_shexp.weight"), - DENSE, - &[shape.ff_expert, shape.embd], - )?; - if layer < shape.hash_layers { - expect( - model, - &name("ffn_gate_tid2eid.weight"), - &[I32], - &[shape.experts_used, shape.vocab], - )?; - } - } - Ok(()) -} - -fn validate_glm_tensors(model: &Gguf, shape: &Shape) -> Result<(), String> { - let q_dim = shape.heads * shape.key_mla; - let q_nope = shape.key_mla - shape.rot; - let index_q = shape.indexer_heads * shape.indexer_head_dim; - expect( - model, - "token_embd.weight", - DENSE, - &[shape.embd, shape.vocab], - )?; - expect(model, "output_norm.weight", &[F32], &[shape.embd])?; - expect(model, "output.weight", DENSE, &[shape.embd, shape.vocab])?; - for layer in 0..shape.layers { - let name = |suffix: &str| format!("blk.{layer}.{suffix}"); - expect(model, &name("attn_norm.weight"), &[F32], &[shape.embd])?; - expect( - model, - &name("attn_q_a.weight"), - DENSE, - &[shape.embd, shape.lora_q], - )?; - expect( - model, - &name("attn_q_a_norm.weight"), - &[F32], - &[shape.lora_q], - )?; - expect( - model, - &name("attn_q_b.weight"), - DENSE, - &[shape.lora_q, q_dim], - )?; - expect( - model, - &name("attn_kv_a_mqa.weight"), - DENSE, - &[shape.embd, shape.head_dim], - )?; - expect( - model, - &name("attn_kv_a_norm.weight"), - &[F32], - &[shape.kv_lora], - )?; - expect( - model, - &name("attn_k_b.weight"), - DENSE, - &[q_nope, shape.kv_lora, shape.heads], - )?; - expect( - model, - &name("attn_v_b.weight"), - DENSE, - &[shape.kv_lora, shape.value_mla, shape.heads], - )?; - expect( - model, - &name("attn_output.weight"), - DENSE, - &[shape.heads * shape.value_mla, shape.embd], - )?; - expect( - model, - &name("indexer.attn_k.weight"), - DENSE, - &[shape.embd, shape.indexer_head_dim], - )?; - expect( - model, - &name("indexer.attn_q_b.weight"), - DENSE, - &[shape.lora_q, index_q], - )?; - expect( - model, - &name("indexer.k_norm.weight"), - &[F32], - &[shape.indexer_head_dim], - )?; - expect( - model, - &name("indexer.k_norm.bias"), - &[F32], - &[shape.indexer_head_dim], - )?; - expect( - model, - &name("indexer.proj.weight"), - &[F32], - &[shape.embd, shape.indexer_heads], - )?; - expect(model, &name("ffn_norm.weight"), &[F32], &[shape.embd])?; - if layer < shape.leading_dense { - expect( - model, - &name("ffn_gate.weight"), - DENSE, - &[shape.embd, shape.ff_dense], - )?; - expect( - model, - &name("ffn_up.weight"), - DENSE, - &[shape.embd, shape.ff_dense], - )?; - expect( - model, - &name("ffn_down.weight"), - DENSE, - &[shape.ff_dense, shape.embd], - )?; - } else { - expect( - model, - &name("ffn_gate_inp.weight"), - &[F32], - &[shape.embd, shape.experts], - )?; - expect(model, &name("exp_probs_b.bias"), &[F32], &[shape.experts])?; - expect( - model, - &name("ffn_gate_exps.weight"), - ROUTED, - &[shape.embd, shape.ff_expert, shape.experts], - )?; - expect( - model, - &name("ffn_up_exps.weight"), - ROUTED, - &[shape.embd, shape.ff_expert, shape.experts], - )?; - expect( - model, - &name("ffn_down_exps.weight"), - ROUTED, - &[shape.ff_expert, shape.embd, shape.experts], - )?; - same_type( - model, - &name("ffn_gate_exps.weight"), - &name("ffn_up_exps.weight"), - )?; - expect( - model, - &name("ffn_gate_shexp.weight"), - DENSE, - &[shape.embd, shape.ff_expert], - )?; - expect( - model, - &name("ffn_up_shexp.weight"), - DENSE, - &[shape.embd, shape.ff_expert], - )?; - expect( - model, - &name("ffn_down_shexp.weight"), - DENSE, - &[shape.ff_expert, shape.embd], - )?; - } - if layer + shape.nextn >= shape.layers { - expect( - model, - &name("nextn.eh_proj.weight"), - DENSE, - &[2 * shape.embd, shape.embd], - )?; - expect(model, &name("nextn.enorm.weight"), &[F32], &[shape.embd])?; - expect(model, &name("nextn.hnorm.weight"), &[F32], &[shape.embd])?; - expect( - model, - &name("nextn.shared_head_norm.weight"), - &[F32], - &[shape.embd], - )?; - } - } - Ok(()) -} - -fn validate_dspark(model: &Gguf, shape: &Shape) -> Result<(), String> { - if shape.model != ModelChoice::DeepSeekV4Flash { - return Err("DSpark support is available only for DeepSeek V4 Flash".into()); - } - let block_size = first_u32( - model, - &[ - "deepseek4.dspark.block_size", - "deepseek4.dspark_block_size", - "dspark.block_size", - ], - )?; - let markov_rank = first_u32( - model, - &[ - "deepseek4.dspark.markov_rank", - "deepseek4.dspark_markov_rank", - "dspark.markov_rank", - ], - )?; - let noise_token = first_u32( - model, - &[ - "deepseek4.dspark.noise_token_id", - "deepseek4.dspark_noise_token_id", - "dspark.noise_token_id", - ], - )?; - let targets = first_u32s( - model, - &[ - "deepseek4.dspark.target_layer_ids", - "deepseek4.dspark_target_layer_ids", - "dspark.target_layer_ids", - ], - )?; - if !(1..=16).contains(&block_size) || markov_rank == 0 || noise_token >= shape.vocab as u32 { - return Err("invalid DSpark block, Markov, or noise-token metadata".into()); - } - if targets.is_empty() - || targets.len() > 8 - || targets.windows(2).any(|pair| pair[0] >= pair[1]) - || targets.iter().any(|layer| *layer >= shape.layers) - { - return Err("invalid DSpark target-layer metadata".into()); - } - let stages = model - .tensors - .keys() - .filter_map(|name| { - name.strip_prefix("mtp.")? - .split('.') - .next()? - .parse::() - .ok() - }) - .max() - .map_or(0, |stage| stage + 1); - if !(1..=8).contains(&stages) { - return Err(format!("invalid DSpark stage count: {stages}")); - } - for stage in 0..stages { - validate_dspark_block(model, shape, stage)?; - if stage == 0 { - expect( - model, - &format!("mtp.{stage}.main_proj.weight"), - DSPARK_DENSE, - &[targets.len() as u64 * shape.embd, shape.embd], - )?; - expect( - model, - &format!("mtp.{stage}.main_norm.weight"), - &[F32], - &[shape.embd], - )?; - } - } - let prefix = format!("mtp.{}", stages - 1); - expect( - model, - &format!("{prefix}.norm.weight"), - &[F32], - &[shape.embd], - )?; - expect( - model, - &format!("{prefix}.hc_head_base.weight"), - &[F32], - &[shape.hc], - )?; - expect( - model, - &format!("{prefix}.hc_head_fn.weight"), - PLAIN, - &[shape.embd * shape.hc, shape.hc], - )?; - expect( - model, - &format!("{prefix}.hc_head_scale.weight"), - &[F32], - &[1], - )?; - expect( - model, - &format!("{prefix}.markov_head.markov_w1.weight"), - DSPARK_DENSE, - &[u64::from(markov_rank), shape.vocab], - )?; - expect( - model, - &format!("{prefix}.markov_head.markov_w2.weight"), - DSPARK_DENSE, - &[u64::from(markov_rank), shape.vocab], - )?; - expect( - model, - &format!("{prefix}.confidence_head.proj.weight"), - DSPARK_DENSE, - &[shape.embd + u64::from(markov_rank), 1], - )?; - Ok(()) -} - -fn validate_dspark_block(model: &Gguf, shape: &Shape, stage: u32) -> Result<(), String> { - let hc_dim = shape.embd * shape.hc; - let hc_mix = 2 * shape.hc + shape.hc * shape.hc; - let q_dim = shape.heads * shape.head_dim; - let output_low = shape.out_groups * shape.lora_o; - let name = |suffix: &str| format!("mtp.{stage}.{suffix}"); - for (suffix, types, dims) in [ - ("hc_attn_fn.weight", PLAIN, vec![hc_dim, hc_mix]), - ("hc_attn_scale.weight", &[F32][..], vec![3]), - ("hc_attn_base.weight", &[F32][..], vec![hc_mix]), - ("attn_norm.weight", &[F32][..], vec![shape.embd]), - ( - "attn_q_a.weight", - DSPARK_DENSE, - vec![shape.embd, shape.lora_q], - ), - ("attn_q_a_norm.weight", &[F32][..], vec![shape.lora_q]), - ("attn_q_b.weight", DSPARK_DENSE, vec![shape.lora_q, q_dim]), - ( - "attn_kv.weight", - DSPARK_DENSE, - vec![shape.embd, shape.head_dim], - ), - ("attn_kv_a_norm.weight", &[F32][..], vec![shape.head_dim]), - ("attn_sinks.weight", &[F32][..], vec![shape.heads]), - ( - "attn_output_a.weight", - DSPARK_DENSE, - vec![ - shape.head_dim * (shape.heads / shape.out_groups), - output_low, - ], - ), - ( - "attn_output_b.weight", - DSPARK_DENSE, - vec![output_low, shape.embd], - ), - ("hc_ffn_fn.weight", PLAIN, vec![hc_dim, hc_mix]), - ("hc_ffn_scale.weight", &[F32][..], vec![3]), - ("hc_ffn_base.weight", &[F32][..], vec![hc_mix]), - ("ffn_norm.weight", &[F32][..], vec![shape.embd]), - ( - "ffn_gate_inp.weight", - DSPARK_DENSE, - vec![shape.embd, shape.experts], - ), - ("exp_probs_b.bias", &[F32][..], vec![shape.experts]), - ( - "ffn_gate_exps.weight", - ROUTED, - vec![shape.embd, shape.ff_expert, shape.experts], - ), - ( - "ffn_up_exps.weight", - ROUTED, - vec![shape.embd, shape.ff_expert, shape.experts], - ), - ( - "ffn_down_exps.weight", - ROUTED, - vec![shape.ff_expert, shape.embd, shape.experts], - ), - ( - "ffn_gate_shexp.weight", - DSPARK_DENSE, - vec![shape.embd, shape.ff_expert], - ), - ( - "ffn_up_shexp.weight", - DSPARK_DENSE, - vec![shape.embd, shape.ff_expert], - ), - ( - "ffn_down_shexp.weight", - DSPARK_DENSE, - vec![shape.ff_expert, shape.embd], - ), - ] { - expect(model, &name(suffix), types, &dims)?; - } - same_type( - model, - &name("ffn_gate_exps.weight"), - &name("ffn_up_exps.weight"), - ) -} - -fn expect(model: &Gguf, name: &str, types: &[u32], dims: &[u64]) -> Result<(), String> { - validate_tensor(name, model.tensor(name)?, types, dims) -} - -fn expect_optional(model: &Gguf, name: &str, types: &[u32], dims: &[u64]) -> Result<(), String> { - match model.tensors.get(name) { - Some(tensor) => validate_tensor(name, tensor, types, dims), - None => Ok(()), - } -} - -fn validate_tensor(name: &str, tensor: &Tensor, types: &[u32], dims: &[u64]) -> Result<(), String> { - if !types.contains(&tensor.kind) { - return Err(format!( - "tensor {name} has unsupported type {}", - tensor.kind - )); - } - if tensor.dims != dims { - return Err(format!( - "tensor {name} has dimensions {:?}, expected {dims:?}", - tensor.dims - )); - } - Ok(()) -} - -fn same_type(model: &Gguf, first: &str, second: &str) -> Result<(), String> { - if model.tensor(first)?.kind != model.tensor(second)?.kind { - Err(format!( - "tensors {first} and {second} use different quantizations" - )) - } else { - Ok(()) - } -} - -fn expect_u64(model: &Gguf, key: &str, expected: u64) -> Result<(), String> { - let actual = model.u64(key)?; - if actual == expected { - Ok(()) - } else { - Err(format!("{key} is {actual}, expected {expected}")) - } -} - -fn expect_float(model: &Gguf, key: &str, expected: f32) -> Result<(), String> { - let actual = model.f32(key)?; - if float_eq(actual, expected) { - Ok(()) - } else { - Err(format!("{key} is {actual}, expected {expected}")) - } -} - -fn float_eq(actual: f32, expected: f32) -> bool { - actual.is_finite() && (actual - expected).abs() <= expected.abs().max(1.0) * 1.0e-6 -} - -fn compression_ratio(shape: &Shape, layer: u32) -> u32 { - match shape.model { - ModelChoice::DeepSeekV4Flash if layer < 2 => 0, - ModelChoice::DeepSeekV4Pro if layer < 2 => 128, - ModelChoice::DeepSeekV4Flash | ModelChoice::DeepSeekV4Pro if layer.is_multiple_of(2) => 4, - ModelChoice::DeepSeekV4Flash | ModelChoice::DeepSeekV4Pro => 128, - ModelChoice::Glm52 => 0, - } -} - -fn first_u32(model: &Gguf, keys: &[&str]) -> Result { - keys.iter() - .find_map(|key| model.u32(key).ok()) - .ok_or_else(|| format!("required DSpark metadata is missing: {}", keys[0])) -} - -fn first_u32s<'a>(model: &'a Gguf, keys: &[&str]) -> Result<&'a [u32], String> { - keys.iter() - .find_map(|key| model.u32s(key).ok()) - .ok_or_else(|| format!("required DSpark metadata is missing: {}", keys[0])) -} - -#[cfg(test)] -mod tests { - use super::*; - - #[test] - fn installed_ds4_fixture_opens_and_renders_a_prompt() { - let path = Path::new( - "../ds4/gguf/DeepSeek-V4-Flash-IQ2XXS-w2Q2K-AProjQ8-SExpQ8-OutQ8-chat-v2-imatrix.gguf", - ); - if !path.exists() { - return; - } - let model = Model::open_main(path, ModelChoice::DeepSeekV4Flash).unwrap(); - let summary = model.summary(); - assert_eq!(summary.model, ModelChoice::DeepSeekV4Flash); - assert_eq!(summary.vocabulary_size, 129_280); - assert_eq!( - model.tokenize("Hello, world! 1234\nint café = 7;\n中文テスト"), - [ - 19_923, 14, 2_058, 3, 223, 6_895, 22, 201, 650, 57_664, 438, 223, 25, 510, 21_134, - 109_288, - ] - ); - assert_eq!( - model.render_prompt("", "Hello", ReasoningMode::Direct), - [0, 128_803, 19_923, 128_804, 128_822] - ); - assert_eq!(model.tokenizer.tokenize_rendered("|DSML|").len(), 1); - assert!(model.is_stop_token_for_reasoning(128_822, ReasoningMode::Direct)); - assert!(!model.is_stop_token_for_reasoning(128_822, ReasoningMode::High)); - let think_start = *model - .render_prompt("", "Hello", ReasoningMode::High) - .last() - .unwrap(); - assert_eq!( - model.render_conversation( - "", - &[ - ChatTurn { - user: true, - skip_previous_eos: false, - reasoning: None, - reasoning_complete: true, - content: "Hello".into(), - }, - ChatTurn { - user: false, - skip_previous_eos: false, - reasoning: None, - reasoning_complete: true, - content: "Hello".into(), - }, - ChatTurn { - user: true, - skip_previous_eos: false, - reasoning: None, - reasoning_complete: true, - content: "Hello".into(), - }, - ], - ReasoningMode::Direct, - ), - [ - 0, 128_803, 19_923, 128_804, 128_822, 19_923, 1, 128_803, 19_923, 128_804, 128_822, - ] - ); - assert_eq!( - model.render_continuation("Hello", ReasoningMode::Direct, false), - [1, 128_803, 19_923, 128_804, 128_822] - ); - assert_eq!( - model.render_continuation("Hello", ReasoningMode::Direct, true), - [128_803, 19_923, 128_804, 128_822] - ); - assert_eq!( - model.render_conversation( - "", - &[ - ChatTurn { - user: true, - skip_previous_eos: false, - reasoning: None, - reasoning_complete: true, - content: "Hello".into(), - }, - ChatTurn { - user: false, - skip_previous_eos: false, - reasoning: Some("Hello".into()), - reasoning_complete: true, - content: "Hello".into(), - }, - ChatTurn { - user: true, - skip_previous_eos: false, - reasoning: None, - reasoning_complete: true, - content: "Hello".into(), - }, - ], - ReasoningMode::Direct, - ), - [ - 0, - 128_803, - 19_923, - 128_804, - think_start, - 19_923, - 128_822, - 19_923, - 1, - 128_803, - 19_923, - 128_804, - 128_822, - ] - ); - } - - #[test] - fn installed_dspark_fixture_passes_the_target_layout() { - let path = Path::new("../ds4/gguf/DeepSeek-V4-Flash-DSpark-support.gguf"); - if path.exists() { - validate_model_artifact(path, ModelChoice::DeepSeekV4Flash, true).unwrap(); - } - } -} diff --git a/src/engine/metal.rs b/src/engine/metal.rs index 809a769..5b0ac2d 100644 --- a/src/engine/metal.rs +++ b/src/engine/metal.rs @@ -1,3 +1,8 @@ +mod checkpoint; +mod gpu; + +use gpu::*; + use super::gguf::{F16, Gguf, Q8_0, Tensor as GgufTensor}; use super::{Model, ModelFamily}; use std::env; @@ -67,923 +72,6 @@ pub(crate) fn configure_sources() -> Result<(), String> { Ok(()) } -#[repr(C)] -struct GpuTensor { - _private: [u8; 0], -} - -unsafe extern "C" { - fn ds4_gpu_init() -> i32; - fn ds4_gpu_cleanup(); - fn ds4_gpu_set_model_map_range( - model_map: *const c_void, - model_size: u64, - map_offset: u64, - map_size: u64, - max_tensor_bytes: u64, - ) -> i32; - fn ds4_gpu_set_quality(quality: bool); - fn ds4_gpu_tensor_alloc(bytes: u64) -> *mut GpuTensor; - fn ds4_gpu_tensor_view(base: *const GpuTensor, offset: u64, bytes: u64) -> *mut GpuTensor; - fn ds4_gpu_tensor_free(tensor: *mut GpuTensor); - fn ds4_gpu_tensor_fill_f32(tensor: *mut GpuTensor, value: f32, count: u64) -> i32; - fn ds4_gpu_tensor_read( - tensor: *const GpuTensor, - offset: u64, - data: *mut c_void, - bytes: u64, - ) -> i32; - fn ds4_gpu_tensor_write( - tensor: *mut GpuTensor, - offset: u64, - data: *const c_void, - bytes: u64, - ) -> i32; - fn ds4_gpu_tensor_copy( - dst: *mut GpuTensor, - dst_offset: u64, - src: *const GpuTensor, - src_offset: u64, - bytes: u64, - ) -> i32; - fn ds4_gpu_tensor_copy_f32_to_f16( - dst: *mut GpuTensor, - dst_offset: u64, - src: *const GpuTensor, - src_offset: u64, - count: u64, - ) -> i32; - fn ds4_gpu_begin_commands() -> i32; - fn ds4_gpu_end_commands() -> i32; - - fn ds4_gpu_embed_tokens_hc_tensor( - out: *mut GpuTensor, - tokens: *const GpuTensor, - map: *const c_void, - size: u64, - weight: u64, - vocab: u32, - rows: u32, - embd: u32, - hc: u32, - ) -> i32; - - fn ds4_gpu_embed_token_hc_tensor( - out: *mut GpuTensor, - map: *const c_void, - size: u64, - weight: u64, - vocab: u32, - token: u32, - embd: u32, - hc: u32, - ) -> i32; - fn ds4_gpu_rms_norm_plain_tensor( - out: *mut GpuTensor, - x: *const GpuTensor, - n: u32, - eps: f32, - ) -> i32; - fn ds4_gpu_rms_norm_weight_tensor( - out: *mut GpuTensor, - x: *const GpuTensor, - map: *const c_void, - size: u64, - weight: u64, - n: u32, - eps: f32, - ) -> i32; - fn ds4_gpu_hc_rms_scale_project_f16_tensor( - out: *mut GpuTensor, - scale: *mut GpuTensor, - map: *const c_void, - size: u64, - weight: u64, - input: u32, - output: u32, - x: *const GpuTensor, - rows: u32, - eps: f32, - ) -> i32; - fn ds4_gpu_matmul_f16_tensor( - out: *mut GpuTensor, - map: *const c_void, - size: u64, - weight: u64, - input: u64, - output: u64, - x: *const GpuTensor, - rows: u64, - ) -> i32; - fn ds4_gpu_matmul_q8_0_tensor( - out: *mut GpuTensor, - map: *const c_void, - size: u64, - weight: u64, - input: u64, - output: u64, - x: *const GpuTensor, - rows: u64, - ) -> i32; - fn ds4_gpu_matmul_q8_0_pair_tensor( - out_a: *mut GpuTensor, - out_b: *mut GpuTensor, - map: *const c_void, - size: u64, - weight_a: u64, - weight_b: u64, - input: u64, - output_a: u64, - output_b: u64, - x: *const GpuTensor, - rows: u64, - ) -> i32; - fn ds4_gpu_matmul_f16_pair_tensor( - out_a: *mut GpuTensor, - out_b: *mut GpuTensor, - map: *const c_void, - size: u64, - weight_a: u64, - weight_b: u64, - input: u64, - output: u64, - x: *const GpuTensor, - rows: u64, - ) -> i32; - fn ds4_gpu_matmul_f16_pair_compressor_store_tensor( - out_kv: *mut GpuTensor, - out_score: *mut GpuTensor, - state_kv: *mut GpuTensor, - state_score: *mut GpuTensor, - map: *const c_void, - size: u64, - weight_kv: u64, - weight_score: u64, - ape: u64, - ape_type: u32, - input: u64, - width: u32, - x: *const GpuTensor, - ratio: u32, - pos: u32, - ) -> i32; - fn ds4_gpu_hc_split_weighted_sum_norm_tensor( - out: *mut GpuTensor, - norm: *mut GpuTensor, - split: *mut GpuTensor, - mix: *const GpuTensor, - residual: *const GpuTensor, - map: *const c_void, - size: u64, - scale: u64, - base: u64, - norm_weight: u64, - embd: u32, - hc: u32, - iterations: u32, - eps: f32, - norm_eps: f32, - ) -> i32; - fn ds4_gpu_dsv4_qkv_rms_norm_rows_tensor( - q_out: *mut GpuTensor, - q: *const GpuTensor, - map: *const c_void, - size: u64, - q_weight: u64, - q_n: u32, - kv_out: *mut GpuTensor, - kv: *const GpuTensor, - kv_weight: u64, - kv_n: u32, - rows: u32, - eps: f32, - ) -> i32; - fn ds4_gpu_attn_q_b_f16_head_rms_rope_tail_tensor( - out: *mut GpuTensor, - half: *mut GpuTensor, - map: *const c_void, - size: u64, - weight: u64, - input: u64, - output: u64, - x: *const GpuTensor, - rows: u32, - heads: u32, - head_dim: u32, - rot: u32, - pos: u32, - original_context: u32, - inverse: bool, - freq_base: f32, - freq_scale: f32, - ext_factor: f32, - attn_factor: f32, - beta_fast: f32, - beta_slow: f32, - eps: f32, - ) -> i32; - fn ds4_gpu_head_rms_norm_tensor( - x: *mut GpuTensor, - rows: u32, - heads: u32, - head_dim: u32, - eps: f32, - ) -> i32; - fn ds4_gpu_head_rms_norm_rope_tail_tensor( - x: *mut GpuTensor, - rows: u32, - heads: u32, - head_dim: u32, - rot: u32, - pos: u32, - original_context: u32, - inverse: bool, - freq_base: f32, - freq_scale: f32, - ext_factor: f32, - attn_factor: f32, - beta_fast: f32, - beta_slow: f32, - eps: f32, - ) -> i32; - fn ds4_gpu_rope_tail_tensor( - x: *mut GpuTensor, - rows: u32, - heads: u32, - head_dim: u32, - rot: u32, - pos: u32, - original_context: u32, - inverse: bool, - freq_base: f32, - freq_scale: f32, - ext_factor: f32, - attn_factor: f32, - beta_fast: f32, - beta_slow: f32, - ) -> i32; - fn ds4_gpu_attention_decode_heads_tensor( - heads_out: *mut GpuTensor, - map: *const c_void, - size: u64, - sinks: u64, - q: *const GpuTensor, - raw_kv: *const GpuTensor, - n_raw: u32, - raw_cap: u32, - raw_start: u32, - comp_kv: *const GpuTensor, - comp_f16: u32, - n_comp: u32, - comp_mask: *const GpuTensor, - use_mask: u32, - heads: u32, - head_dim: u32, - ) -> i32; - fn ds4_gpu_attention_decode_raw_batch_heads_tensor( - heads: *mut GpuTensor, - map: *const c_void, - size: u64, - sinks: u64, - q: *const GpuTensor, - raw_kv: *const GpuTensor, - tokens: u32, - pos: u32, - n_raw: u32, - raw_cap: u32, - raw_start: u32, - window: u32, - heads_count: u32, - head_dim: u32, - ) -> i32; - fn ds4_gpu_attention_decode_mixed_batch_heads_tensor( - heads: *mut GpuTensor, - map: *const c_void, - size: u64, - sinks: u64, - q: *const GpuTensor, - raw_kv: *const GpuTensor, - compressed: *const GpuTensor, - compressed_f16: u32, - compressed_mask: *const GpuTensor, - use_mask: u32, - tokens: u32, - pos: u32, - n_raw: u32, - raw_cap: u32, - raw_start: u32, - n_comp: u32, - window: u32, - ratio: u32, - heads_count: u32, - head_dim: u32, - ) -> i32; - fn ds4_gpu_attention_indexed_mixed_batch_heads_tensor( - heads: *mut GpuTensor, - map: *const c_void, - size: u64, - sinks: u64, - q: *const GpuTensor, - raw_kv: *const GpuTensor, - comp_kv: *const GpuTensor, - comp_f16: u32, - topk: *const GpuTensor, - tokens: u32, - pos: u32, - n_raw: u32, - raw_cap: u32, - raw_start: u32, - n_comp: u32, - top_k: u32, - window: u32, - ratio: u32, - heads: u32, - head_dim: u32, - ) -> i32; - fn ds4_gpu_indexer_score_one_tensor( - scores: *mut GpuTensor, - q: *const GpuTensor, - weights: *const GpuTensor, - index_comp: *const GpuTensor, - n_comp: u32, - heads: u32, - head_dim: u32, - scale: f32, - ) -> i32; - fn ds4_gpu_indexer_scores_decode_batch_tensor( - scores: *mut GpuTensor, - q: *const GpuTensor, - weights: *const GpuTensor, - index_comp: *const GpuTensor, - n_comp: u32, - tokens: u32, - pos: u32, - heads: u32, - head_dim: u32, - ratio: u32, - scale: f32, - ) -> i32; - fn ds4_gpu_indexer_topk_tensor( - selected: *mut GpuTensor, - scores: *const GpuTensor, - n_comp: u32, - tokens: u32, - top_k: u32, - ) -> i32; - fn ds4_gpu_dsv4_indexer_qat_tensor(x: *mut GpuTensor, rows: u32, head_dim: u32) -> i32; - fn ds4_gpu_compressor_update_tensor( - kv: *const GpuTensor, - score: *const GpuTensor, - state_kv: *mut GpuTensor, - state_score: *mut GpuTensor, - cache: *mut GpuTensor, - map: *const c_void, - size: u64, - ape: u64, - ape_type: u32, - norm: u64, - norm_type: u32, - head_dim: u32, - ratio: u32, - pos: u32, - row: u32, - rot: u32, - original_context: u32, - freq_base: f32, - freq_scale: f32, - ext_factor: f32, - attn_factor: f32, - beta_fast: f32, - beta_slow: f32, - rms_eps: f32, - state_already_stored: bool, - ) -> i32; - fn ds4_gpu_compressor_prefill_state_ratio4_tensor( - state_kv: *mut GpuTensor, - state_score: *mut GpuTensor, - kv_tail: *const GpuTensor, - score_tail: *const GpuTensor, - map: *const c_void, - size: u64, - ape: u64, - ape_type: u32, - head_dim: u32, - pos: u32, - ) -> i32; - fn ds4_gpu_dsv4_fp8_kv_quantize_tensor( - x: *mut GpuTensor, - rows: u32, - head_dim: u32, - rot: u32, - ) -> i32; - fn ds4_gpu_kv_fp8_store_raw_tensor( - kv: *mut GpuTensor, - raw_cache: *mut GpuTensor, - raw_cap: u32, - row: u32, - head_dim: u32, - rot: u32, - ) -> i32; - fn ds4_gpu_store_raw_kv_batch_tensor( - raw_cache: *mut GpuTensor, - kv: *const GpuTensor, - raw_cap: u32, - pos: u32, - rows: u32, - head_dim: u32, - ) -> i32; - fn ds4_gpu_compressor_prefill_tensor( - cache: *mut GpuTensor, - state_kv: *mut GpuTensor, - state_score: *mut GpuTensor, - kv: *const GpuTensor, - score: *const GpuTensor, - map: *const c_void, - size: u64, - ape: u64, - ape_type: u32, - norm: u64, - norm_type: u32, - head_dim: u32, - ratio: u32, - pos: u32, - rows: u32, - rot: u32, - original_context: u32, - quantize_fp8: bool, - freq_base: f32, - freq_scale: f32, - ext_factor: f32, - attn_factor: f32, - beta_fast: f32, - beta_slow: f32, - rms_eps: f32, - ) -> i32; - fn ds4_gpu_compressor_prefill_ratio4_replay_tensor( - cache: *mut GpuTensor, - state_kv: *mut GpuTensor, - state_score: *mut GpuTensor, - kv: *const GpuTensor, - score: *const GpuTensor, - map: *const c_void, - size: u64, - ape: u64, - ape_type: u32, - norm: u64, - norm_type: u32, - head_dim: u32, - pos: u32, - rows: u32, - rot: u32, - original_context: u32, - quantize_fp8: bool, - freq_base: f32, - freq_scale: f32, - ext_factor: f32, - attn_factor: f32, - beta_fast: f32, - beta_slow: f32, - rms_eps: f32, - ) -> i32; - fn ds4_gpu_attention_prefill_raw_heads_tensor( - heads: *mut GpuTensor, - map: *const c_void, - size: u64, - sinks: u64, - q: *const GpuTensor, - raw: *const GpuTensor, - rows: u32, - window: u32, - heads_count: u32, - head_dim: u32, - ) -> i32; - fn ds4_gpu_attention_prefill_static_mixed_heads_tensor( - heads: *mut GpuTensor, - map: *const c_void, - size: u64, - sinks: u64, - q: *const GpuTensor, - raw: *const GpuTensor, - compressed: *const GpuTensor, - compressed_f16: u32, - rows: u32, - compressed_rows: u32, - window: u32, - ratio: u32, - heads_count: u32, - head_dim: u32, - ) -> i32; - fn ds4_gpu_indexer_scores_prefill_tensor( - scores: *mut GpuTensor, - q: *const GpuTensor, - weights: *const GpuTensor, - compressed: *const GpuTensor, - compressed_rows: u32, - rows: u32, - heads: u32, - head_dim: u32, - ratio: u32, - scale: f32, - ) -> i32; - fn ds4_gpu_attention_output_q8_batch_f16_tensor( - out_half: *mut GpuTensor, - low: *mut GpuTensor, - map: *const c_void, - size: u64, - weight_a: u64, - weight_b: u64, - group_dim: u64, - rank: u64, - groups: u32, - output: u64, - heads: *const GpuTensor, - rows: u32, - ) -> i32; - fn ds4_gpu_attention_output_q8_batch_tensor( - out: *mut GpuTensor, - low: *mut GpuTensor, - group_scratch: *mut GpuTensor, - low_scratch: *mut GpuTensor, - map: *const c_void, - size: u64, - weight_a: u64, - weight_b: u64, - group_dim: u64, - rank: u64, - groups: u32, - output: u64, - heads: *const GpuTensor, - rows: u32, - ) -> i32; - fn ds4_gpu_router_select_batch_tensor( - selected: *mut GpuTensor, - weights: *mut GpuTensor, - probs: *mut GpuTensor, - map: *const c_void, - size: u64, - bias: u64, - hash: u64, - hash_rows: u32, - expert_groups: u32, - groups_used: u32, - has_bias: bool, - hash_mode: bool, - logits: *const GpuTensor, - tokens: *const GpuTensor, - experts: u32, - experts_used: u32, - scale: f32, - rows: u32, - ) -> i32; - fn ds4_gpu_routed_moe_batch_tensor( - out: *mut GpuTensor, - gate: *mut GpuTensor, - up: *mut GpuTensor, - mid: *mut GpuTensor, - experts_out: *mut GpuTensor, - map: *const c_void, - size: u64, - gate_weight: u64, - up_weight: u64, - down_weight: u64, - gate_type: u32, - down_type: u32, - gate_expert_bytes: u64, - gate_row_bytes: u64, - down_expert_bytes: u64, - down_row_bytes: u64, - input: u32, - middle: u32, - output: u32, - selected: *const GpuTensor, - weights: *const GpuTensor, - total_experts: u32, - used_experts: u32, - clamp: f32, - x: *const GpuTensor, - layer: u32, - rows: u32, - mid_f16: *mut bool, - force_resident: bool, - ) -> i32; - fn ds4_gpu_swiglu_tensor( - out: *mut GpuTensor, - gate: *const GpuTensor, - up: *const GpuTensor, - count: u32, - clamp: f32, - scale: f32, - ) -> i32; - fn ds4_gpu_hc_expand_split_half_tensor( - out: *mut GpuTensor, - block_half: *const GpuTensor, - residual: *const GpuTensor, - split: *const GpuTensor, - embd: u32, - hc: u32, - ) -> i32; - fn ds4_gpu_hc_expand_split_tensor( - out: *mut GpuTensor, - block: *const GpuTensor, - residual: *const GpuTensor, - split: *const GpuTensor, - embd: u32, - hc: u32, - ) -> i32; - fn ds4_gpu_hc_expand_add_split_tensor( - out: *mut GpuTensor, - block: *const GpuTensor, - add: *const GpuTensor, - residual: *const GpuTensor, - split: *const GpuTensor, - embd: u32, - hc: u32, - ) -> i32; - fn ds4_gpu_attention_output_low_q8_tensor( - low: *mut GpuTensor, - map: *const c_void, - size: u64, - weight: u64, - group_dim: u64, - rank: u64, - groups: u32, - heads: *const GpuTensor, - ) -> i32; - fn ds4_gpu_matmul_q8_0_hc_expand_tensor( - out_hc: *mut GpuTensor, - block_out: *mut GpuTensor, - map: *const c_void, - size: u64, - weight: u64, - input: u64, - output: u64, - x: *const GpuTensor, - residual: *const GpuTensor, - split: *const GpuTensor, - embd: u32, - hc: u32, - ) -> i32; - fn ds4_gpu_router_select_tensor( - selected: *mut GpuTensor, - weights: *mut GpuTensor, - probs: *mut GpuTensor, - map: *const c_void, - size: u64, - bias: u64, - hash: u64, - hash_rows: u32, - token: u32, - experts: u32, - used: u32, - scale: f32, - expert_groups: u32, - groups_used: u32, - has_bias: bool, - hash_mode: bool, - logits: *const GpuTensor, - ) -> i32; - fn ds4_gpu_routed_moe_one_tensor( - out: *mut GpuTensor, - gate: *mut GpuTensor, - up: *mut GpuTensor, - mid: *mut GpuTensor, - experts_out: *mut GpuTensor, - map: *const c_void, - size: u64, - gate_weight: u64, - up_weight: u64, - down_weight: u64, - gate_type: u32, - down_type: u32, - gate_expert_bytes: u64, - gate_row_bytes: u64, - down_expert_bytes: u64, - down_row_bytes: u64, - input: u32, - middle: u32, - output: u32, - selected: *const GpuTensor, - weights: *const GpuTensor, - total_experts: u32, - used_experts: u32, - clamp: f32, - x: *const GpuTensor, - add: *const GpuTensor, - layer: u32, - force_resident: bool, - ) -> i32; - fn ds4_gpu_shared_gate_up_swiglu_q8_0_tensor( - gate: *mut GpuTensor, - up: *mut GpuTensor, - mid: *mut GpuTensor, - map: *const c_void, - size: u64, - gate_weight: u64, - up_weight: u64, - input: u64, - output: u64, - x: *const GpuTensor, - clamp: f32, - ) -> i32; - fn ds4_gpu_shared_down_hc_expand_q8_0_tensor( - out_hc: *mut GpuTensor, - shared_out: *mut GpuTensor, - map: *const c_void, - size: u64, - weight: u64, - input: u64, - output: u64, - middle: *const GpuTensor, - routed: *const GpuTensor, - residual: *const GpuTensor, - split: *const GpuTensor, - embd: u32, - hc: u32, - ) -> i32; - fn ds4_gpu_output_hc_weights_tensor( - out: *mut GpuTensor, - pre: *const GpuTensor, - map: *const c_void, - size: u64, - scale: u64, - base: u64, - hc: u32, - eps: f32, - ) -> i32; - fn ds4_gpu_hc_weighted_sum_norm_tensor( - out: *mut GpuTensor, - norm: *mut GpuTensor, - residual: *const GpuTensor, - weights: *const GpuTensor, - map: *const c_void, - size: u64, - norm_weight: u64, - embd: u32, - hc: u32, - eps: f32, - ) -> i32; - fn ds4_gpu_hc_weighted_sum_tensor( - out: *mut GpuTensor, - residual: *const GpuTensor, - weights: *const GpuTensor, - embd: u32, - hc: u32, - ) -> i32; -} - -struct Context; - -impl Context { - fn open(model: &Model, quality: bool) -> Result { - check(unsafe { ds4_gpu_init() }, "Metal initialization")?; - let data_offset = model.main.data_offset(); - if let Err(error) = check( - unsafe { - ds4_gpu_set_model_map_range( - model.main.map_ptr().cast(), - model.main.len(), - data_offset, - model.main.len() - data_offset, - model.main.max_tensor_bytes(), - ) - }, - "model mapping", - ) { - unsafe { ds4_gpu_cleanup() }; - return Err(error); - } - unsafe { ds4_gpu_set_quality(quality) }; - Ok(Self) - } -} - -impl Drop for Context { - fn drop(&mut self) { - unsafe { ds4_gpu_cleanup() }; - } -} - -struct Commands(bool); - -impl Commands { - fn begin() -> Result { - check( - unsafe { ds4_gpu_begin_commands() }, - "beginning Metal commands", - )?; - Ok(Self(true)) - } - - fn finish(mut self) -> Result<(), String> { - self.0 = false; - check( - unsafe { ds4_gpu_end_commands() }, - "executing Metal commands", - ) - } -} - -impl Drop for Commands { - fn drop(&mut self) { - if self.0 { - unsafe { ds4_gpu_end_commands() }; - } - } -} - -struct Buffer(NonNull); - -impl Buffer { - fn floats(count: u64) -> Result { - Self::bytes(count * 4) - } - - fn bytes(bytes: u64) -> Result { - NonNull::new(unsafe { ds4_gpu_tensor_alloc(bytes) }) - .map(Self) - .ok_or_else(|| format!("Metal could not allocate {bytes} bytes")) - } - - fn view(&self, offset: u64, bytes: u64) -> Result { - NonNull::new(unsafe { ds4_gpu_tensor_view(self.raw(), offset, bytes) }) - .map(Self) - .ok_or_else(|| "Metal could not create a tensor view".to_owned()) - } - - fn read_f32(&self, values: &mut [f32]) -> Result<(), String> { - check( - unsafe { - ds4_gpu_tensor_read( - self.raw(), - 0, - values.as_mut_ptr().cast(), - std::mem::size_of_val(values) as u64, - ) - }, - "reading Metal output", - ) - } - - fn read(&self, offset: u64, values: &mut [u8]) -> Result<(), String> { - check( - unsafe { - ds4_gpu_tensor_read( - self.raw(), - offset, - values.as_mut_ptr().cast(), - values.len() as u64, - ) - }, - "reading a Metal buffer", - ) - } - - fn write(&self, offset: u64, values: &[u8]) -> Result<(), String> { - check( - unsafe { - ds4_gpu_tensor_write( - self.raw(), - offset, - values.as_ptr().cast(), - values.len() as u64, - ) - }, - "restoring a Metal buffer", - ) - } - - fn write_i32(&self, values: &[i32]) -> Result<(), String> { - check( - unsafe { - ds4_gpu_tensor_write( - self.raw(), - 0, - values.as_ptr().cast(), - std::mem::size_of_val(values) as u64, - ) - }, - "uploading tokens", - ) - } - - fn fill(&self, value: f32, count: u64) -> Result<(), String> { - call( - unsafe { ds4_gpu_tensor_fill_f32(self.raw(), value, count) }, - "initializing a Metal buffer", - ) - } - - fn raw(&self) -> *mut GpuTensor { - self.0.as_ptr() - } -} - -impl Drop for Buffer { - fn drop(&mut self) { - unsafe { ds4_gpu_tensor_free(self.raw()) }; - } -} - #[derive(Clone, Copy)] struct Weight { offset: u64, @@ -1708,301 +796,6 @@ impl Executor { self.checkpoint_tag } - pub(super) fn save_checkpoint( - &mut self, - path: &Path, - tag: [u8; 32], - progress: &mut impl FnMut(u64), - ) -> Result<(), String> { - if let Some(parent) = path.parent() { - fs::create_dir_all(parent).map_err(|error| error.to_string())?; - } - let temporary = path.with_extension("tmp"); - let mut file = File::create(&temporary).map_err(|error| error.to_string())?; - let mut reported = 0_u64; - let mut pending = 0_u64; - self.write_checkpoint(&mut file, tag, &mut |bytes| { - reported += bytes; - pending += bytes; - if pending >= CHECKPOINT_IO_CHUNK as u64 { - progress(pending); - pending = 0; - } - })?; - progress(pending); - let total = file.metadata().map_err(|error| error.to_string())?.len(); - progress(total.saturating_sub(reported)); - file.sync_all().map_err(|error| error.to_string())?; - fs::rename(&temporary, path).map_err(|error| error.to_string())?; - self.checkpoint_tag = tag; - Ok(()) - } - - pub(super) fn load_checkpoint( - &mut self, - path: &Path, - progress: &mut impl FnMut(u64), - ) -> Result { - let mut file = match File::open(path) { - Ok(file) => file, - Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(false), - Err(error) => return Err(error.to_string()), - }; - let mut reported = 0_u64; - let mut pending = 0_u64; - self.read_checkpoint(&mut file, &mut |bytes| { - reported += bytes; - pending += bytes; - if pending >= CHECKPOINT_IO_CHUNK as u64 { - progress(pending); - pending = 0; - } - })?; - progress(pending); - let total = file.metadata().map_err(|error| error.to_string())?.len(); - progress(total.saturating_sub(reported)); - Ok(true) - } - - fn write_checkpoint( - &self, - file: &mut File, - tag: [u8; 32], - progress: &mut impl FnMut(u64), - ) -> Result<(), String> { - let shape = self.model.shape; - file.write_all(CHECKPOINT_MAGIC) - .map_err(|error| error.to_string())?; - for value in [ - CHECKPOINT_VERSION, - self.session.context, - self.session.raw_cap, - shape.layers, - shape.head_dim as u32, - shape.indexer_head_dim as u32, - shape.vocab as u32, - u32::from(self.quality), - ] { - write_u32(file, value)?; - } - write_u64(file, self.model.main.len())?; - write_u64(file, self.model_modified.0)?; - write_u32(file, self.model_modified.1)?; - for weight in [self.weights.token_embedding, self.weights.output] { - write_u64(file, weight.offset)?; - write_u64(file, weight.bytes)?; - write_u32(file, weight.kind)?; - } - file.write_all(&tag).map_err(|error| error.to_string())?; - let token_count = u32::try_from(self.tokens.len()) - .map_err(|_| "KV checkpoint has too many tokens".to_owned())?; - let raw_live = token_count.min(self.session.raw_cap); - write_u32(file, token_count)?; - write_u32(file, raw_live)?; - for &token in &self.tokens { - write_u32(file, token as u32)?; - } - for &logit in &self.logits { - write_u32(file, logit.to_bits())?; - } - for layer in &self.session.layers { - write_u32( - file, - layer.compression.as_ref().map_or(0, |state| state.rows), - )?; - } - for layer in &self.session.layers { - write_u32(file, layer.indexer.as_ref().map_or(0, |state| state.rows))?; - } - - let mut chunk = vec![0; CHECKPOINT_IO_CHUNK]; - let raw_first = token_count - raw_live; - for layer in &self.session.layers { - for position in raw_first..token_count { - let physical = position % self.session.raw_cap; - write_buffer( - file, - &layer.raw_cache, - u64::from(physical) * shape.head_dim * 4, - shape.head_dim * 4, - &mut chunk, - progress, - )?; - } - if let Some(state) = &layer.compression { - write_buffer( - file, - &state.cache, - 0, - u64::from(state.rows) * shape.head_dim * 2, - &mut chunk, - progress, - )?; - let bytes = compressor_state_bytes(state.ratio, shape.head_dim); - write_buffer(file, &state.state_kv, 0, bytes, &mut chunk, progress)?; - write_buffer(file, &state.state_score, 0, bytes, &mut chunk, progress)?; - } - if let Some(state) = &layer.indexer { - write_buffer( - file, - &state.cache, - 0, - u64::from(state.rows) * shape.indexer_head_dim * 4, - &mut chunk, - progress, - )?; - let bytes = compressor_state_bytes(state.ratio, shape.indexer_head_dim); - write_buffer(file, &state.state_kv, 0, bytes, &mut chunk, progress)?; - write_buffer(file, &state.state_score, 0, bytes, &mut chunk, progress)?; - } - } - Ok(()) - } - - fn read_checkpoint( - &mut self, - file: &mut File, - progress: &mut impl FnMut(u64), - ) -> Result<(), String> { - let mut magic = [0; 8]; - file.read_exact(&mut magic) - .map_err(|error| error.to_string())?; - if &magic != CHECKPOINT_MAGIC { - return Err("KV checkpoint has an invalid signature".into()); - } - let shape = self.model.shape; - let header = [ - CHECKPOINT_VERSION, - self.session.context, - self.session.raw_cap, - shape.layers, - shape.head_dim as u32, - shape.indexer_head_dim as u32, - shape.vocab as u32, - u32::from(self.quality), - ]; - for expected in header { - if read_u32(file)? != expected { - return Err("KV checkpoint does not match the current executor".into()); - } - } - if read_u64(file)? != self.model.main.len() { - return Err("KV checkpoint was written for a different model".into()); - } - if read_u64(file)? != self.model_modified.0 || read_u32(file)? != self.model_modified.1 { - return Err("KV checkpoint model file has changed".into()); - } - for weight in [self.weights.token_embedding, self.weights.output] { - if read_u64(file)? != weight.offset - || read_u64(file)? != weight.bytes - || read_u32(file)? != weight.kind - { - return Err("KV checkpoint was written for a different model layout".into()); - } - } - let mut checkpoint_tag = [0; 32]; - file.read_exact(&mut checkpoint_tag) - .map_err(|error| error.to_string())?; - let token_count = read_u32(file)?; - let raw_live = read_u32(file)?; - if token_count > self.session.context || raw_live != token_count.min(self.session.raw_cap) { - return Err("KV checkpoint token count is invalid".into()); - } - let mut tokens = Vec::with_capacity(token_count as usize); - for _ in 0..token_count { - let token = read_u32(file)?; - if u64::from(token) >= shape.vocab { - return Err("KV checkpoint contains an invalid token".into()); - } - tokens.push(token as i32); - } - let mut logits = Vec::with_capacity(shape.vocab as usize); - for _ in 0..shape.vocab { - logits.push(f32::from_bits(read_u32(file)?)); - } - let mut compressed_rows = Vec::with_capacity(shape.layers as usize); - let mut indexer_rows = Vec::with_capacity(shape.layers as usize); - for layer in 0..shape.layers { - let rows = read_u32(file)?; - let ratio = compression_ratio(layer); - if rows != token_count.checked_div(ratio).unwrap_or(0) { - return Err("KV checkpoint compressed row count is invalid".into()); - } - compressed_rows.push(rows); - } - for layer in 0..shape.layers { - let rows = read_u32(file)?; - let expected = if compression_ratio(layer) == 4 { - token_count / 4 - } else { - 0 - }; - if rows != expected { - return Err("KV checkpoint indexer row count is invalid".into()); - } - indexer_rows.push(rows); - } - - self.reset()?; - let mut chunk = vec![0; CHECKPOINT_IO_CHUNK]; - let raw_first = token_count - raw_live; - for (index, layer) in self.session.layers.iter_mut().enumerate() { - for position in raw_first..token_count { - let physical = position % self.session.raw_cap; - read_buffer( - file, - &layer.raw_cache, - u64::from(physical) * shape.head_dim * 4, - shape.head_dim * 4, - &mut chunk, - progress, - )?; - } - if let Some(state) = &mut layer.compression { - state.rows = compressed_rows[index]; - read_buffer( - file, - &state.cache, - 0, - u64::from(state.rows) * shape.head_dim * 2, - &mut chunk, - progress, - )?; - let bytes = compressor_state_bytes(state.ratio, shape.head_dim); - read_buffer(file, &state.state_kv, 0, bytes, &mut chunk, progress)?; - read_buffer(file, &state.state_score, 0, bytes, &mut chunk, progress)?; - } - if let Some(state) = &mut layer.indexer { - state.rows = indexer_rows[index]; - read_buffer( - file, - &state.cache, - 0, - u64::from(state.rows) * shape.indexer_head_dim * 4, - &mut chunk, - progress, - )?; - let bytes = compressor_state_bytes(state.ratio, shape.indexer_head_dim); - read_buffer(file, &state.state_kv, 0, bytes, &mut chunk, progress)?; - read_buffer(file, &state.state_score, 0, bytes, &mut chunk, progress)?; - } - } - let mut trailing = [0]; - if file - .read(&mut trailing) - .map_err(|error| error.to_string())? - != 0 - { - self.reset()?; - return Err("KV checkpoint has trailing data".into()); - } - self.session.position = token_count; - self.tokens = tokens; - self.logits = logits; - self.checkpoint_tag = checkpoint_tag; - Ok(()) - } - fn encode_token(&mut self, token: u32) -> Result<(), String> { let shape = self.model.shape; let map = self.model.main.map_ptr().cast(); @@ -4173,75 +2966,6 @@ fn raw_batch_span(pos: u32, rows: u32, raw_cap: u32, window: u32) -> (u32, u32) (needed, (last + 1 - needed) % raw_cap) } -fn compressor_state_bytes(ratio: u32, head_dim: u64) -> u64 { - let coefficient = if ratio == 4 { 2 } else { 1 }; - coefficient * head_dim * coefficient * u64::from(ratio) * 4 -} - -fn write_buffer( - file: &mut File, - buffer: &Buffer, - mut offset: u64, - mut bytes: u64, - chunk: &mut [u8], - progress: &mut impl FnMut(u64), -) -> Result<(), String> { - while bytes != 0 { - let length = bytes.min(chunk.len() as u64) as usize; - buffer.read(offset, &mut chunk[..length])?; - file.write_all(&chunk[..length]) - .map_err(|error| error.to_string())?; - progress(length as u64); - offset += length as u64; - bytes -= length as u64; - } - Ok(()) -} - -fn read_buffer( - file: &mut File, - buffer: &Buffer, - mut offset: u64, - mut bytes: u64, - chunk: &mut [u8], - progress: &mut impl FnMut(u64), -) -> Result<(), String> { - while bytes != 0 { - let length = bytes.min(chunk.len() as u64) as usize; - file.read_exact(&mut chunk[..length]) - .map_err(|error| error.to_string())?; - buffer.write(offset, &chunk[..length])?; - progress(length as u64); - offset += length as u64; - bytes -= length as u64; - } - Ok(()) -} - -fn write_u32(file: &mut File, value: u32) -> Result<(), String> { - file.write_all(&value.to_le_bytes()) - .map_err(|error| error.to_string()) -} - -fn write_u64(file: &mut File, value: u64) -> Result<(), String> { - file.write_all(&value.to_le_bytes()) - .map_err(|error| error.to_string()) -} - -fn read_u32(file: &mut File) -> Result { - let mut bytes = [0; 4]; - file.read_exact(&mut bytes) - .map_err(|error| error.to_string())?; - Ok(u32::from_le_bytes(bytes)) -} - -fn read_u64(file: &mut File) -> Result { - let mut bytes = [0; 8]; - file.read_exact(&mut bytes) - .map_err(|error| error.to_string())?; - Ok(u64::from_le_bytes(bytes)) -} - fn encode_output( s: &Scratch, w: &Weights, diff --git a/src/engine/metal/checkpoint.rs b/src/engine/metal/checkpoint.rs new file mode 100644 index 0000000..945d42e --- /dev/null +++ b/src/engine/metal/checkpoint.rs @@ -0,0 +1,367 @@ +use super::*; + +impl Executor { + pub(in crate::engine) fn save_checkpoint( + &mut self, + path: &Path, + tag: [u8; 32], + progress: &mut impl FnMut(u64), + ) -> Result<(), String> { + if let Some(parent) = path.parent() { + fs::create_dir_all(parent).map_err(|error| error.to_string())?; + } + let temporary = path.with_extension("tmp"); + let mut file = File::create(&temporary).map_err(|error| error.to_string())?; + let mut reported = 0_u64; + let mut pending = 0_u64; + self.write_checkpoint(&mut file, tag, &mut |bytes| { + reported += bytes; + pending += bytes; + if pending >= CHECKPOINT_IO_CHUNK as u64 { + progress(pending); + pending = 0; + } + })?; + progress(pending); + let total = file.metadata().map_err(|error| error.to_string())?.len(); + progress(total.saturating_sub(reported)); + file.sync_all().map_err(|error| error.to_string())?; + fs::rename(&temporary, path).map_err(|error| error.to_string())?; + self.checkpoint_tag = tag; + Ok(()) + } + + pub(in crate::engine) fn load_checkpoint( + &mut self, + path: &Path, + progress: &mut impl FnMut(u64), + ) -> Result { + let mut file = match File::open(path) { + Ok(file) => file, + Err(error) if error.kind() == std::io::ErrorKind::NotFound => return Ok(false), + Err(error) => return Err(error.to_string()), + }; + let mut reported = 0_u64; + let mut pending = 0_u64; + self.read_checkpoint(&mut file, &mut |bytes| { + reported += bytes; + pending += bytes; + if pending >= CHECKPOINT_IO_CHUNK as u64 { + progress(pending); + pending = 0; + } + })?; + progress(pending); + let total = file.metadata().map_err(|error| error.to_string())?.len(); + progress(total.saturating_sub(reported)); + Ok(true) + } + + fn write_checkpoint( + &self, + file: &mut File, + tag: [u8; 32], + progress: &mut impl FnMut(u64), + ) -> Result<(), String> { + let shape = self.model.shape; + file.write_all(CHECKPOINT_MAGIC) + .map_err(|error| error.to_string())?; + for value in [ + CHECKPOINT_VERSION, + self.session.context, + self.session.raw_cap, + shape.layers, + shape.head_dim as u32, + shape.indexer_head_dim as u32, + shape.vocab as u32, + u32::from(self.quality), + ] { + write_u32(file, value)?; + } + write_u64(file, self.model.main.len())?; + write_u64(file, self.model_modified.0)?; + write_u32(file, self.model_modified.1)?; + for weight in [self.weights.token_embedding, self.weights.output] { + write_u64(file, weight.offset)?; + write_u64(file, weight.bytes)?; + write_u32(file, weight.kind)?; + } + file.write_all(&tag).map_err(|error| error.to_string())?; + let token_count = u32::try_from(self.tokens.len()) + .map_err(|_| "KV checkpoint has too many tokens".to_owned())?; + let raw_live = token_count.min(self.session.raw_cap); + write_u32(file, token_count)?; + write_u32(file, raw_live)?; + for &token in &self.tokens { + write_u32(file, token as u32)?; + } + for &logit in &self.logits { + write_u32(file, logit.to_bits())?; + } + for layer in &self.session.layers { + write_u32( + file, + layer.compression.as_ref().map_or(0, |state| state.rows), + )?; + } + for layer in &self.session.layers { + write_u32(file, layer.indexer.as_ref().map_or(0, |state| state.rows))?; + } + + let mut chunk = vec![0; CHECKPOINT_IO_CHUNK]; + let raw_first = token_count - raw_live; + for layer in &self.session.layers { + for position in raw_first..token_count { + let physical = position % self.session.raw_cap; + write_buffer( + file, + &layer.raw_cache, + u64::from(physical) * shape.head_dim * 4, + shape.head_dim * 4, + &mut chunk, + progress, + )?; + } + if let Some(state) = &layer.compression { + write_buffer( + file, + &state.cache, + 0, + u64::from(state.rows) * shape.head_dim * 2, + &mut chunk, + progress, + )?; + let bytes = compressor_state_bytes(state.ratio, shape.head_dim); + write_buffer(file, &state.state_kv, 0, bytes, &mut chunk, progress)?; + write_buffer(file, &state.state_score, 0, bytes, &mut chunk, progress)?; + } + if let Some(state) = &layer.indexer { + write_buffer( + file, + &state.cache, + 0, + u64::from(state.rows) * shape.indexer_head_dim * 4, + &mut chunk, + progress, + )?; + let bytes = compressor_state_bytes(state.ratio, shape.indexer_head_dim); + write_buffer(file, &state.state_kv, 0, bytes, &mut chunk, progress)?; + write_buffer(file, &state.state_score, 0, bytes, &mut chunk, progress)?; + } + } + Ok(()) + } + + fn read_checkpoint( + &mut self, + file: &mut File, + progress: &mut impl FnMut(u64), + ) -> Result<(), String> { + let mut magic = [0; 8]; + file.read_exact(&mut magic) + .map_err(|error| error.to_string())?; + if &magic != CHECKPOINT_MAGIC { + return Err("KV checkpoint has an invalid signature".into()); + } + let shape = self.model.shape; + let header = [ + CHECKPOINT_VERSION, + self.session.context, + self.session.raw_cap, + shape.layers, + shape.head_dim as u32, + shape.indexer_head_dim as u32, + shape.vocab as u32, + u32::from(self.quality), + ]; + for expected in header { + if read_u32(file)? != expected { + return Err("KV checkpoint does not match the current executor".into()); + } + } + if read_u64(file)? != self.model.main.len() { + return Err("KV checkpoint was written for a different model".into()); + } + if read_u64(file)? != self.model_modified.0 || read_u32(file)? != self.model_modified.1 { + return Err("KV checkpoint model file has changed".into()); + } + for weight in [self.weights.token_embedding, self.weights.output] { + if read_u64(file)? != weight.offset + || read_u64(file)? != weight.bytes + || read_u32(file)? != weight.kind + { + return Err("KV checkpoint was written for a different model layout".into()); + } + } + let mut checkpoint_tag = [0; 32]; + file.read_exact(&mut checkpoint_tag) + .map_err(|error| error.to_string())?; + let token_count = read_u32(file)?; + let raw_live = read_u32(file)?; + if token_count > self.session.context || raw_live != token_count.min(self.session.raw_cap) { + return Err("KV checkpoint token count is invalid".into()); + } + let mut tokens = Vec::with_capacity(token_count as usize); + for _ in 0..token_count { + let token = read_u32(file)?; + if u64::from(token) >= shape.vocab { + return Err("KV checkpoint contains an invalid token".into()); + } + tokens.push(token as i32); + } + let mut logits = Vec::with_capacity(shape.vocab as usize); + for _ in 0..shape.vocab { + logits.push(f32::from_bits(read_u32(file)?)); + } + let mut compressed_rows = Vec::with_capacity(shape.layers as usize); + let mut indexer_rows = Vec::with_capacity(shape.layers as usize); + for layer in 0..shape.layers { + let rows = read_u32(file)?; + let ratio = compression_ratio(layer); + if rows != token_count.checked_div(ratio).unwrap_or(0) { + return Err("KV checkpoint compressed row count is invalid".into()); + } + compressed_rows.push(rows); + } + for layer in 0..shape.layers { + let rows = read_u32(file)?; + let expected = if compression_ratio(layer) == 4 { + token_count / 4 + } else { + 0 + }; + if rows != expected { + return Err("KV checkpoint indexer row count is invalid".into()); + } + indexer_rows.push(rows); + } + + self.reset()?; + let mut chunk = vec![0; CHECKPOINT_IO_CHUNK]; + let raw_first = token_count - raw_live; + for (index, layer) in self.session.layers.iter_mut().enumerate() { + for position in raw_first..token_count { + let physical = position % self.session.raw_cap; + read_buffer( + file, + &layer.raw_cache, + u64::from(physical) * shape.head_dim * 4, + shape.head_dim * 4, + &mut chunk, + progress, + )?; + } + if let Some(state) = &mut layer.compression { + state.rows = compressed_rows[index]; + read_buffer( + file, + &state.cache, + 0, + u64::from(state.rows) * shape.head_dim * 2, + &mut chunk, + progress, + )?; + let bytes = compressor_state_bytes(state.ratio, shape.head_dim); + read_buffer(file, &state.state_kv, 0, bytes, &mut chunk, progress)?; + read_buffer(file, &state.state_score, 0, bytes, &mut chunk, progress)?; + } + if let Some(state) = &mut layer.indexer { + state.rows = indexer_rows[index]; + read_buffer( + file, + &state.cache, + 0, + u64::from(state.rows) * shape.indexer_head_dim * 4, + &mut chunk, + progress, + )?; + let bytes = compressor_state_bytes(state.ratio, shape.indexer_head_dim); + read_buffer(file, &state.state_kv, 0, bytes, &mut chunk, progress)?; + read_buffer(file, &state.state_score, 0, bytes, &mut chunk, progress)?; + } + } + let mut trailing = [0]; + if file + .read(&mut trailing) + .map_err(|error| error.to_string())? + != 0 + { + self.reset()?; + return Err("KV checkpoint has trailing data".into()); + } + self.session.position = token_count; + self.tokens = tokens; + self.logits = logits; + self.checkpoint_tag = checkpoint_tag; + Ok(()) + } +} + +fn compressor_state_bytes(ratio: u32, head_dim: u64) -> u64 { + let coefficient = if ratio == 4 { 2 } else { 1 }; + coefficient * head_dim * coefficient * u64::from(ratio) * 4 +} + +fn write_buffer( + file: &mut File, + buffer: &Buffer, + mut offset: u64, + mut bytes: u64, + chunk: &mut [u8], + progress: &mut impl FnMut(u64), +) -> Result<(), String> { + while bytes != 0 { + let length = bytes.min(chunk.len() as u64) as usize; + buffer.read(offset, &mut chunk[..length])?; + file.write_all(&chunk[..length]) + .map_err(|error| error.to_string())?; + progress(length as u64); + offset += length as u64; + bytes -= length as u64; + } + Ok(()) +} + +fn read_buffer( + file: &mut File, + buffer: &Buffer, + mut offset: u64, + mut bytes: u64, + chunk: &mut [u8], + progress: &mut impl FnMut(u64), +) -> Result<(), String> { + while bytes != 0 { + let length = bytes.min(chunk.len() as u64) as usize; + file.read_exact(&mut chunk[..length]) + .map_err(|error| error.to_string())?; + buffer.write(offset, &chunk[..length])?; + progress(length as u64); + offset += length as u64; + bytes -= length as u64; + } + Ok(()) +} + +fn write_u32(file: &mut File, value: u32) -> Result<(), String> { + file.write_all(&value.to_le_bytes()) + .map_err(|error| error.to_string()) +} + +fn write_u64(file: &mut File, value: u64) -> Result<(), String> { + file.write_all(&value.to_le_bytes()) + .map_err(|error| error.to_string()) +} + +fn read_u32(file: &mut File) -> Result { + let mut bytes = [0; 4]; + file.read_exact(&mut bytes) + .map_err(|error| error.to_string())?; + Ok(u32::from_le_bytes(bytes)) +} + +fn read_u64(file: &mut File) -> Result { + let mut bytes = [0; 8]; + file.read_exact(&mut bytes) + .map_err(|error| error.to_string())?; + Ok(u64::from_le_bytes(bytes)) +} diff --git a/src/engine/metal/gpu.rs b/src/engine/metal/gpu.rs new file mode 100644 index 0000000..a473659 --- /dev/null +++ b/src/engine/metal/gpu.rs @@ -0,0 +1,926 @@ +use super::*; + +#[repr(C)] +pub(super) struct GpuTensor { + _private: [u8; 0], +} + +unsafe extern "C" { + pub(super) fn ds4_gpu_init() -> i32; + pub(super) fn ds4_gpu_cleanup(); + pub(super) fn ds4_gpu_set_model_map_range( + model_map: *const c_void, + model_size: u64, + map_offset: u64, + map_size: u64, + max_tensor_bytes: u64, + ) -> i32; + pub(super) fn ds4_gpu_set_quality(quality: bool); + pub(super) fn ds4_gpu_tensor_alloc(bytes: u64) -> *mut GpuTensor; + pub(super) fn ds4_gpu_tensor_view( + base: *const GpuTensor, + offset: u64, + bytes: u64, + ) -> *mut GpuTensor; + pub(super) fn ds4_gpu_tensor_free(tensor: *mut GpuTensor); + pub(super) fn ds4_gpu_tensor_fill_f32(tensor: *mut GpuTensor, value: f32, count: u64) -> i32; + pub(super) fn ds4_gpu_tensor_read( + tensor: *const GpuTensor, + offset: u64, + data: *mut c_void, + bytes: u64, + ) -> i32; + pub(super) fn ds4_gpu_tensor_write( + tensor: *mut GpuTensor, + offset: u64, + data: *const c_void, + bytes: u64, + ) -> i32; + pub(super) fn ds4_gpu_tensor_copy( + dst: *mut GpuTensor, + dst_offset: u64, + src: *const GpuTensor, + src_offset: u64, + bytes: u64, + ) -> i32; + pub(super) fn ds4_gpu_tensor_copy_f32_to_f16( + dst: *mut GpuTensor, + dst_offset: u64, + src: *const GpuTensor, + src_offset: u64, + count: u64, + ) -> i32; + pub(super) fn ds4_gpu_begin_commands() -> i32; + pub(super) fn ds4_gpu_end_commands() -> i32; + + pub(super) fn ds4_gpu_embed_tokens_hc_tensor( + out: *mut GpuTensor, + tokens: *const GpuTensor, + map: *const c_void, + size: u64, + weight: u64, + vocab: u32, + rows: u32, + embd: u32, + hc: u32, + ) -> i32; + + pub(super) fn ds4_gpu_embed_token_hc_tensor( + out: *mut GpuTensor, + map: *const c_void, + size: u64, + weight: u64, + vocab: u32, + token: u32, + embd: u32, + hc: u32, + ) -> i32; + pub(super) fn ds4_gpu_rms_norm_plain_tensor( + out: *mut GpuTensor, + x: *const GpuTensor, + n: u32, + eps: f32, + ) -> i32; + pub(super) fn ds4_gpu_rms_norm_weight_tensor( + out: *mut GpuTensor, + x: *const GpuTensor, + map: *const c_void, + size: u64, + weight: u64, + n: u32, + eps: f32, + ) -> i32; + pub(super) fn ds4_gpu_hc_rms_scale_project_f16_tensor( + out: *mut GpuTensor, + scale: *mut GpuTensor, + map: *const c_void, + size: u64, + weight: u64, + input: u32, + output: u32, + x: *const GpuTensor, + rows: u32, + eps: f32, + ) -> i32; + pub(super) fn ds4_gpu_matmul_f16_tensor( + out: *mut GpuTensor, + map: *const c_void, + size: u64, + weight: u64, + input: u64, + output: u64, + x: *const GpuTensor, + rows: u64, + ) -> i32; + pub(super) fn ds4_gpu_matmul_q8_0_tensor( + out: *mut GpuTensor, + map: *const c_void, + size: u64, + weight: u64, + input: u64, + output: u64, + x: *const GpuTensor, + rows: u64, + ) -> i32; + pub(super) fn ds4_gpu_matmul_q8_0_pair_tensor( + out_a: *mut GpuTensor, + out_b: *mut GpuTensor, + map: *const c_void, + size: u64, + weight_a: u64, + weight_b: u64, + input: u64, + output_a: u64, + output_b: u64, + x: *const GpuTensor, + rows: u64, + ) -> i32; + pub(super) fn ds4_gpu_matmul_f16_pair_tensor( + out_a: *mut GpuTensor, + out_b: *mut GpuTensor, + map: *const c_void, + size: u64, + weight_a: u64, + weight_b: u64, + input: u64, + output: u64, + x: *const GpuTensor, + rows: u64, + ) -> i32; + pub(super) fn ds4_gpu_matmul_f16_pair_compressor_store_tensor( + out_kv: *mut GpuTensor, + out_score: *mut GpuTensor, + state_kv: *mut GpuTensor, + state_score: *mut GpuTensor, + map: *const c_void, + size: u64, + weight_kv: u64, + weight_score: u64, + ape: u64, + ape_type: u32, + input: u64, + width: u32, + x: *const GpuTensor, + ratio: u32, + pos: u32, + ) -> i32; + pub(super) fn ds4_gpu_hc_split_weighted_sum_norm_tensor( + out: *mut GpuTensor, + norm: *mut GpuTensor, + split: *mut GpuTensor, + mix: *const GpuTensor, + residual: *const GpuTensor, + map: *const c_void, + size: u64, + scale: u64, + base: u64, + norm_weight: u64, + embd: u32, + hc: u32, + iterations: u32, + eps: f32, + norm_eps: f32, + ) -> i32; + pub(super) fn ds4_gpu_dsv4_qkv_rms_norm_rows_tensor( + q_out: *mut GpuTensor, + q: *const GpuTensor, + map: *const c_void, + size: u64, + q_weight: u64, + q_n: u32, + kv_out: *mut GpuTensor, + kv: *const GpuTensor, + kv_weight: u64, + kv_n: u32, + rows: u32, + eps: f32, + ) -> i32; + pub(super) fn ds4_gpu_attn_q_b_f16_head_rms_rope_tail_tensor( + out: *mut GpuTensor, + half: *mut GpuTensor, + map: *const c_void, + size: u64, + weight: u64, + input: u64, + output: u64, + x: *const GpuTensor, + rows: u32, + heads: u32, + head_dim: u32, + rot: u32, + pos: u32, + original_context: u32, + inverse: bool, + freq_base: f32, + freq_scale: f32, + ext_factor: f32, + attn_factor: f32, + beta_fast: f32, + beta_slow: f32, + eps: f32, + ) -> i32; + pub(super) fn ds4_gpu_head_rms_norm_tensor( + x: *mut GpuTensor, + rows: u32, + heads: u32, + head_dim: u32, + eps: f32, + ) -> i32; + pub(super) fn ds4_gpu_head_rms_norm_rope_tail_tensor( + x: *mut GpuTensor, + rows: u32, + heads: u32, + head_dim: u32, + rot: u32, + pos: u32, + original_context: u32, + inverse: bool, + freq_base: f32, + freq_scale: f32, + ext_factor: f32, + attn_factor: f32, + beta_fast: f32, + beta_slow: f32, + eps: f32, + ) -> i32; + pub(super) fn ds4_gpu_rope_tail_tensor( + x: *mut GpuTensor, + rows: u32, + heads: u32, + head_dim: u32, + rot: u32, + pos: u32, + original_context: u32, + inverse: bool, + freq_base: f32, + freq_scale: f32, + ext_factor: f32, + attn_factor: f32, + beta_fast: f32, + beta_slow: f32, + ) -> i32; + pub(super) fn ds4_gpu_attention_decode_heads_tensor( + heads_out: *mut GpuTensor, + map: *const c_void, + size: u64, + sinks: u64, + q: *const GpuTensor, + raw_kv: *const GpuTensor, + n_raw: u32, + raw_cap: u32, + raw_start: u32, + comp_kv: *const GpuTensor, + comp_f16: u32, + n_comp: u32, + comp_mask: *const GpuTensor, + use_mask: u32, + heads: u32, + head_dim: u32, + ) -> i32; + pub(super) fn ds4_gpu_attention_decode_raw_batch_heads_tensor( + heads: *mut GpuTensor, + map: *const c_void, + size: u64, + sinks: u64, + q: *const GpuTensor, + raw_kv: *const GpuTensor, + tokens: u32, + pos: u32, + n_raw: u32, + raw_cap: u32, + raw_start: u32, + window: u32, + heads_count: u32, + head_dim: u32, + ) -> i32; + pub(super) fn ds4_gpu_attention_decode_mixed_batch_heads_tensor( + heads: *mut GpuTensor, + map: *const c_void, + size: u64, + sinks: u64, + q: *const GpuTensor, + raw_kv: *const GpuTensor, + compressed: *const GpuTensor, + compressed_f16: u32, + compressed_mask: *const GpuTensor, + use_mask: u32, + tokens: u32, + pos: u32, + n_raw: u32, + raw_cap: u32, + raw_start: u32, + n_comp: u32, + window: u32, + ratio: u32, + heads_count: u32, + head_dim: u32, + ) -> i32; + pub(super) fn ds4_gpu_attention_indexed_mixed_batch_heads_tensor( + heads: *mut GpuTensor, + map: *const c_void, + size: u64, + sinks: u64, + q: *const GpuTensor, + raw_kv: *const GpuTensor, + comp_kv: *const GpuTensor, + comp_f16: u32, + topk: *const GpuTensor, + tokens: u32, + pos: u32, + n_raw: u32, + raw_cap: u32, + raw_start: u32, + n_comp: u32, + top_k: u32, + window: u32, + ratio: u32, + heads: u32, + head_dim: u32, + ) -> i32; + pub(super) fn ds4_gpu_indexer_score_one_tensor( + scores: *mut GpuTensor, + q: *const GpuTensor, + weights: *const GpuTensor, + index_comp: *const GpuTensor, + n_comp: u32, + heads: u32, + head_dim: u32, + scale: f32, + ) -> i32; + pub(super) fn ds4_gpu_indexer_scores_decode_batch_tensor( + scores: *mut GpuTensor, + q: *const GpuTensor, + weights: *const GpuTensor, + index_comp: *const GpuTensor, + n_comp: u32, + tokens: u32, + pos: u32, + heads: u32, + head_dim: u32, + ratio: u32, + scale: f32, + ) -> i32; + pub(super) fn ds4_gpu_indexer_topk_tensor( + selected: *mut GpuTensor, + scores: *const GpuTensor, + n_comp: u32, + tokens: u32, + top_k: u32, + ) -> i32; + pub(super) fn ds4_gpu_dsv4_indexer_qat_tensor( + x: *mut GpuTensor, + rows: u32, + head_dim: u32, + ) -> i32; + pub(super) fn ds4_gpu_compressor_update_tensor( + kv: *const GpuTensor, + score: *const GpuTensor, + state_kv: *mut GpuTensor, + state_score: *mut GpuTensor, + cache: *mut GpuTensor, + map: *const c_void, + size: u64, + ape: u64, + ape_type: u32, + norm: u64, + norm_type: u32, + head_dim: u32, + ratio: u32, + pos: u32, + row: u32, + rot: u32, + original_context: u32, + freq_base: f32, + freq_scale: f32, + ext_factor: f32, + attn_factor: f32, + beta_fast: f32, + beta_slow: f32, + rms_eps: f32, + state_already_stored: bool, + ) -> i32; + pub(super) fn ds4_gpu_compressor_prefill_state_ratio4_tensor( + state_kv: *mut GpuTensor, + state_score: *mut GpuTensor, + kv_tail: *const GpuTensor, + score_tail: *const GpuTensor, + map: *const c_void, + size: u64, + ape: u64, + ape_type: u32, + head_dim: u32, + pos: u32, + ) -> i32; + pub(super) fn ds4_gpu_dsv4_fp8_kv_quantize_tensor( + x: *mut GpuTensor, + rows: u32, + head_dim: u32, + rot: u32, + ) -> i32; + pub(super) fn ds4_gpu_kv_fp8_store_raw_tensor( + kv: *mut GpuTensor, + raw_cache: *mut GpuTensor, + raw_cap: u32, + row: u32, + head_dim: u32, + rot: u32, + ) -> i32; + pub(super) fn ds4_gpu_store_raw_kv_batch_tensor( + raw_cache: *mut GpuTensor, + kv: *const GpuTensor, + raw_cap: u32, + pos: u32, + rows: u32, + head_dim: u32, + ) -> i32; + pub(super) fn ds4_gpu_compressor_prefill_tensor( + cache: *mut GpuTensor, + state_kv: *mut GpuTensor, + state_score: *mut GpuTensor, + kv: *const GpuTensor, + score: *const GpuTensor, + map: *const c_void, + size: u64, + ape: u64, + ape_type: u32, + norm: u64, + norm_type: u32, + head_dim: u32, + ratio: u32, + pos: u32, + rows: u32, + rot: u32, + original_context: u32, + quantize_fp8: bool, + freq_base: f32, + freq_scale: f32, + ext_factor: f32, + attn_factor: f32, + beta_fast: f32, + beta_slow: f32, + rms_eps: f32, + ) -> i32; + pub(super) fn ds4_gpu_compressor_prefill_ratio4_replay_tensor( + cache: *mut GpuTensor, + state_kv: *mut GpuTensor, + state_score: *mut GpuTensor, + kv: *const GpuTensor, + score: *const GpuTensor, + map: *const c_void, + size: u64, + ape: u64, + ape_type: u32, + norm: u64, + norm_type: u32, + head_dim: u32, + pos: u32, + rows: u32, + rot: u32, + original_context: u32, + quantize_fp8: bool, + freq_base: f32, + freq_scale: f32, + ext_factor: f32, + attn_factor: f32, + beta_fast: f32, + beta_slow: f32, + rms_eps: f32, + ) -> i32; + pub(super) fn ds4_gpu_attention_prefill_raw_heads_tensor( + heads: *mut GpuTensor, + map: *const c_void, + size: u64, + sinks: u64, + q: *const GpuTensor, + raw: *const GpuTensor, + rows: u32, + window: u32, + heads_count: u32, + head_dim: u32, + ) -> i32; + pub(super) fn ds4_gpu_attention_prefill_static_mixed_heads_tensor( + heads: *mut GpuTensor, + map: *const c_void, + size: u64, + sinks: u64, + q: *const GpuTensor, + raw: *const GpuTensor, + compressed: *const GpuTensor, + compressed_f16: u32, + rows: u32, + compressed_rows: u32, + window: u32, + ratio: u32, + heads_count: u32, + head_dim: u32, + ) -> i32; + pub(super) fn ds4_gpu_indexer_scores_prefill_tensor( + scores: *mut GpuTensor, + q: *const GpuTensor, + weights: *const GpuTensor, + compressed: *const GpuTensor, + compressed_rows: u32, + rows: u32, + heads: u32, + head_dim: u32, + ratio: u32, + scale: f32, + ) -> i32; + pub(super) fn ds4_gpu_attention_output_q8_batch_f16_tensor( + out_half: *mut GpuTensor, + low: *mut GpuTensor, + map: *const c_void, + size: u64, + weight_a: u64, + weight_b: u64, + group_dim: u64, + rank: u64, + groups: u32, + output: u64, + heads: *const GpuTensor, + rows: u32, + ) -> i32; + pub(super) fn ds4_gpu_attention_output_q8_batch_tensor( + out: *mut GpuTensor, + low: *mut GpuTensor, + group_scratch: *mut GpuTensor, + low_scratch: *mut GpuTensor, + map: *const c_void, + size: u64, + weight_a: u64, + weight_b: u64, + group_dim: u64, + rank: u64, + groups: u32, + output: u64, + heads: *const GpuTensor, + rows: u32, + ) -> i32; + pub(super) fn ds4_gpu_router_select_batch_tensor( + selected: *mut GpuTensor, + weights: *mut GpuTensor, + probs: *mut GpuTensor, + map: *const c_void, + size: u64, + bias: u64, + hash: u64, + hash_rows: u32, + expert_groups: u32, + groups_used: u32, + has_bias: bool, + hash_mode: bool, + logits: *const GpuTensor, + tokens: *const GpuTensor, + experts: u32, + experts_used: u32, + scale: f32, + rows: u32, + ) -> i32; + pub(super) fn ds4_gpu_routed_moe_batch_tensor( + out: *mut GpuTensor, + gate: *mut GpuTensor, + up: *mut GpuTensor, + mid: *mut GpuTensor, + experts_out: *mut GpuTensor, + map: *const c_void, + size: u64, + gate_weight: u64, + up_weight: u64, + down_weight: u64, + gate_type: u32, + down_type: u32, + gate_expert_bytes: u64, + gate_row_bytes: u64, + down_expert_bytes: u64, + down_row_bytes: u64, + input: u32, + middle: u32, + output: u32, + selected: *const GpuTensor, + weights: *const GpuTensor, + total_experts: u32, + used_experts: u32, + clamp: f32, + x: *const GpuTensor, + layer: u32, + rows: u32, + mid_f16: *mut bool, + force_resident: bool, + ) -> i32; + pub(super) fn ds4_gpu_swiglu_tensor( + out: *mut GpuTensor, + gate: *const GpuTensor, + up: *const GpuTensor, + count: u32, + clamp: f32, + scale: f32, + ) -> i32; + pub(super) fn ds4_gpu_hc_expand_split_half_tensor( + out: *mut GpuTensor, + block_half: *const GpuTensor, + residual: *const GpuTensor, + split: *const GpuTensor, + embd: u32, + hc: u32, + ) -> i32; + pub(super) fn ds4_gpu_hc_expand_split_tensor( + out: *mut GpuTensor, + block: *const GpuTensor, + residual: *const GpuTensor, + split: *const GpuTensor, + embd: u32, + hc: u32, + ) -> i32; + pub(super) fn ds4_gpu_hc_expand_add_split_tensor( + out: *mut GpuTensor, + block: *const GpuTensor, + add: *const GpuTensor, + residual: *const GpuTensor, + split: *const GpuTensor, + embd: u32, + hc: u32, + ) -> i32; + pub(super) fn ds4_gpu_attention_output_low_q8_tensor( + low: *mut GpuTensor, + map: *const c_void, + size: u64, + weight: u64, + group_dim: u64, + rank: u64, + groups: u32, + heads: *const GpuTensor, + ) -> i32; + pub(super) fn ds4_gpu_matmul_q8_0_hc_expand_tensor( + out_hc: *mut GpuTensor, + block_out: *mut GpuTensor, + map: *const c_void, + size: u64, + weight: u64, + input: u64, + output: u64, + x: *const GpuTensor, + residual: *const GpuTensor, + split: *const GpuTensor, + embd: u32, + hc: u32, + ) -> i32; + pub(super) fn ds4_gpu_router_select_tensor( + selected: *mut GpuTensor, + weights: *mut GpuTensor, + probs: *mut GpuTensor, + map: *const c_void, + size: u64, + bias: u64, + hash: u64, + hash_rows: u32, + token: u32, + experts: u32, + used: u32, + scale: f32, + expert_groups: u32, + groups_used: u32, + has_bias: bool, + hash_mode: bool, + logits: *const GpuTensor, + ) -> i32; + pub(super) fn ds4_gpu_routed_moe_one_tensor( + out: *mut GpuTensor, + gate: *mut GpuTensor, + up: *mut GpuTensor, + mid: *mut GpuTensor, + experts_out: *mut GpuTensor, + map: *const c_void, + size: u64, + gate_weight: u64, + up_weight: u64, + down_weight: u64, + gate_type: u32, + down_type: u32, + gate_expert_bytes: u64, + gate_row_bytes: u64, + down_expert_bytes: u64, + down_row_bytes: u64, + input: u32, + middle: u32, + output: u32, + selected: *const GpuTensor, + weights: *const GpuTensor, + total_experts: u32, + used_experts: u32, + clamp: f32, + x: *const GpuTensor, + add: *const GpuTensor, + layer: u32, + force_resident: bool, + ) -> i32; + pub(super) fn ds4_gpu_shared_gate_up_swiglu_q8_0_tensor( + gate: *mut GpuTensor, + up: *mut GpuTensor, + mid: *mut GpuTensor, + map: *const c_void, + size: u64, + gate_weight: u64, + up_weight: u64, + input: u64, + output: u64, + x: *const GpuTensor, + clamp: f32, + ) -> i32; + pub(super) fn ds4_gpu_shared_down_hc_expand_q8_0_tensor( + out_hc: *mut GpuTensor, + shared_out: *mut GpuTensor, + map: *const c_void, + size: u64, + weight: u64, + input: u64, + output: u64, + middle: *const GpuTensor, + routed: *const GpuTensor, + residual: *const GpuTensor, + split: *const GpuTensor, + embd: u32, + hc: u32, + ) -> i32; + pub(super) fn ds4_gpu_output_hc_weights_tensor( + out: *mut GpuTensor, + pre: *const GpuTensor, + map: *const c_void, + size: u64, + scale: u64, + base: u64, + hc: u32, + eps: f32, + ) -> i32; + pub(super) fn ds4_gpu_hc_weighted_sum_norm_tensor( + out: *mut GpuTensor, + norm: *mut GpuTensor, + residual: *const GpuTensor, + weights: *const GpuTensor, + map: *const c_void, + size: u64, + norm_weight: u64, + embd: u32, + hc: u32, + eps: f32, + ) -> i32; + pub(super) fn ds4_gpu_hc_weighted_sum_tensor( + out: *mut GpuTensor, + residual: *const GpuTensor, + weights: *const GpuTensor, + embd: u32, + hc: u32, + ) -> i32; +} + +pub(super) struct Context; + +impl Context { + pub(super) fn open(model: &Model, quality: bool) -> Result { + check(unsafe { ds4_gpu_init() }, "Metal initialization")?; + let data_offset = model.main.data_offset(); + if let Err(error) = check( + unsafe { + ds4_gpu_set_model_map_range( + model.main.map_ptr().cast(), + model.main.len(), + data_offset, + model.main.len() - data_offset, + model.main.max_tensor_bytes(), + ) + }, + "model mapping", + ) { + unsafe { ds4_gpu_cleanup() }; + return Err(error); + } + unsafe { ds4_gpu_set_quality(quality) }; + Ok(Self) + } +} + +impl Drop for Context { + fn drop(&mut self) { + unsafe { ds4_gpu_cleanup() }; + } +} + +pub(super) struct Commands(bool); + +impl Commands { + pub(super) fn begin() -> Result { + check( + unsafe { ds4_gpu_begin_commands() }, + "beginning Metal commands", + )?; + Ok(Self(true)) + } + + pub(super) fn finish(mut self) -> Result<(), String> { + self.0 = false; + check( + unsafe { ds4_gpu_end_commands() }, + "executing Metal commands", + ) + } +} + +impl Drop for Commands { + fn drop(&mut self) { + if self.0 { + unsafe { ds4_gpu_end_commands() }; + } + } +} + +pub(super) struct Buffer(NonNull); + +impl Buffer { + pub(super) fn floats(count: u64) -> Result { + Self::bytes(count * 4) + } + + pub(super) fn bytes(bytes: u64) -> Result { + NonNull::new(unsafe { ds4_gpu_tensor_alloc(bytes) }) + .map(Self) + .ok_or_else(|| format!("Metal could not allocate {bytes} bytes")) + } + + pub(super) fn view(&self, offset: u64, bytes: u64) -> Result { + NonNull::new(unsafe { ds4_gpu_tensor_view(self.raw(), offset, bytes) }) + .map(Self) + .ok_or_else(|| "Metal could not create a tensor view".to_owned()) + } + + pub(super) fn read_f32(&self, values: &mut [f32]) -> Result<(), String> { + check( + unsafe { + ds4_gpu_tensor_read( + self.raw(), + 0, + values.as_mut_ptr().cast(), + std::mem::size_of_val(values) as u64, + ) + }, + "reading Metal output", + ) + } + + pub(super) fn read(&self, offset: u64, values: &mut [u8]) -> Result<(), String> { + check( + unsafe { + ds4_gpu_tensor_read( + self.raw(), + offset, + values.as_mut_ptr().cast(), + values.len() as u64, + ) + }, + "reading a Metal buffer", + ) + } + + pub(super) fn write(&self, offset: u64, values: &[u8]) -> Result<(), String> { + check( + unsafe { + ds4_gpu_tensor_write( + self.raw(), + offset, + values.as_ptr().cast(), + values.len() as u64, + ) + }, + "restoring a Metal buffer", + ) + } + + pub(super) fn write_i32(&self, values: &[i32]) -> Result<(), String> { + check( + unsafe { + ds4_gpu_tensor_write( + self.raw(), + 0, + values.as_ptr().cast(), + std::mem::size_of_val(values) as u64, + ) + }, + "uploading tokens", + ) + } + + pub(super) fn fill(&self, value: f32, count: u64) -> Result<(), String> { + call( + unsafe { ds4_gpu_tensor_fill_f32(self.raw(), value, count) }, + "initializing a Metal buffer", + ) + } + + pub(super) fn raw(&self) -> *mut GpuTensor { + self.0.as_ptr() + } +} + +impl Drop for Buffer { + fn drop(&mut self) { + unsafe { ds4_gpu_tensor_free(self.raw()) }; + } +} diff --git a/src/engine/validation.rs b/src/engine/validation.rs new file mode 100644 index 0000000..42c0eab --- /dev/null +++ b/src/engine/validation.rs @@ -0,0 +1,994 @@ +use super::*; + +pub(crate) fn validate_model_artifact( + path: &Path, + expected: ModelChoice, + support: bool, +) -> Result<(), String> { + if support { + let model = Gguf::open(path)?; + validate_dspark(&model, &FLASH) + } else { + let model = Model::open_main(path, expected)?; + let summary = model.summary(); + if summary.tensor_count == 0 + || model + .render_prompt("", "", ReasoningMode::Direct) + .is_empty() + { + return Err("model intake produced an empty tensor directory or prompt".into()); + } + let _ = model.tokenize(""); + let eos = model.eos_token(); + let _ = model.token_bytes(eos); + if !model.is_stop_token(eos) { + return Err("tokenizer EOS marker is not a stop token".into()); + } + model.tensor_data("token_embd.weight")?; + Ok(()) + } +} + +pub(super) fn validate_main(model: &Gguf, expected: ModelChoice) -> Result { + let family = if model.bytes("general.architecture").ok() == Some(b"glm-dsa") { + ModelFamily::Glm + } else { + ModelFamily::DeepSeek + }; + let shape = match family { + ModelFamily::Glm => GLM, + ModelFamily::DeepSeek => match model.u32("deepseek4.block_count")? { + 43 => FLASH, + 61 => PRO, + layers => return Err(format!("unsupported DeepSeek layer count: {layers}")), + }, + }; + if shape.model != expected { + return Err(format!( + "{} contains {}, but the selected model is {expected}", + model.path().display(), + shape.model + )); + } + validate_metadata(model, &shape)?; + validate_tensors(model, &shape)?; + Ok(shape) +} + +fn validate_metadata(model: &Gguf, shape: &Shape) -> Result<(), String> { + let prefix = if shape.family == ModelFamily::Glm { + "glm-dsa" + } else { + "deepseek4" + }; + for (key, expected) in [ + ("block_count", u64::from(shape.layers)), + ("embedding_length", shape.embd), + ("vocab_size", shape.vocab), + ("attention.head_count", shape.heads), + ("attention.head_count_kv", shape.head_kv), + ("attention.key_length", shape.head_dim), + ("attention.value_length", shape.value_dim), + ("rope.dimension_count", shape.rot), + ("attention.q_lora_rank", shape.lora_q), + ("expert_count", shape.experts), + ("expert_used_count", shape.experts_used), + ("expert_feed_forward_length", shape.ff_expert), + ("expert_shared_count", shape.expert_shared), + ("attention.indexer.head_count", shape.indexer_heads), + ("attention.indexer.key_length", shape.indexer_head_dim), + ("attention.indexer.top_k", shape.indexer_top_k), + ] { + expect_u64(model, &format!("{prefix}.{key}"), expected)?; + } + expect_float( + model, + &format!("{prefix}.attention.layer_norm_rms_epsilon"), + shape.rms_epsilon, + )?; + expect_float( + model, + &format!("{prefix}.expert_weights_scale"), + shape.expert_weight_scale, + )?; + if !model.boolean(&format!("{prefix}.expert_weights_norm"))? { + return Err(format!("{prefix}.expert_weights_norm must be true")); + } + expect_float(model, &format!("{prefix}.rope.freq_base"), shape.rope_base)?; + + if shape.family == ModelFamily::Glm { + for (key, expected) in [ + ("context_length", shape.original_context), + ("feed_forward_length", shape.ff_dense), + ("attention.kv_lora_rank", shape.kv_lora), + ("attention.key_length_mla", shape.key_mla), + ("attention.value_length_mla", shape.value_mla), + ("expert_group_count", 1), + ("expert_group_used_count", 1), + ("expert_gating_func", 2), + ("leading_dense_block_count", u64::from(shape.leading_dense)), + ("nextn_predict_layers", u64::from(shape.nextn)), + ] { + expect_u64(model, &format!("{prefix}.{key}"), expected)?; + } + return Ok(()); + } + + for (key, expected) in [ + ("attention.output_group_count", shape.out_groups), + ("attention.output_lora_rank", shape.lora_o), + ("hash_layer_count", u64::from(shape.hash_layers)), + ("attention.sliding_window", shape.sliding_window), + ("hyper_connection.count", shape.hc), + ("hyper_connection.sinkhorn_iterations", shape.hc_sinkhorn), + ] { + expect_u64(model, &format!("{prefix}.{key}"), expected)?; + } + for key in ["expert_group_count", "expert_group_used_count"] { + if let Some(Value::U32(value)) = model.metadata.get(&format!("{prefix}.{key}")) + && *value != 0 + { + return Err(format!("{prefix}.{key} must be zero")); + } + } + for (key, expected) in [ + ("hyper_connection.epsilon", shape.hc_epsilon), + ( + "attention.compress_rope_freq_base", + shape.compress_rope_base, + ), + ] { + expect_float(model, &format!("{prefix}.{key}"), expected)?; + } + for (key, expected) in [ + ("rope.scaling.factor", shape.rope_scale), + ("rope.scaling.yarn_beta_fast", shape.rope_beta_fast), + ("rope.scaling.yarn_beta_slow", shape.rope_beta_slow), + ] { + if model.metadata.contains_key(&format!("{prefix}.{key}")) { + expect_float(model, &format!("{prefix}.{key}"), expected)?; + } + } + if model + .metadata + .contains_key("deepseek4.rope.scaling.original_context_length") + { + expect_u64( + model, + "deepseek4.rope.scaling.original_context_length", + shape.original_context, + )?; + } + let ratios = model.u32s("deepseek4.attention.compress_ratios")?; + let clamps = model.f32s("deepseek4.swiglu_clamp_exp")?; + if ratios.len() < shape.layers as usize || clamps.len() < shape.layers as usize { + return Err("DeepSeek per-layer metadata is shorter than the layer count".into()); + } + for layer in 0..shape.layers as usize { + let expected = compression_ratio(shape, layer as u32); + if ratios[layer] != expected { + return Err(format!( + "layer {layer} compression ratio is {}, expected {expected}", + ratios[layer] + )); + } + if !float_eq(clamps[layer], shape.swiglu_clamp) { + return Err(format!("layer {layer} has an invalid SwiGLU clamp")); + } + } + Ok(()) +} + +fn validate_tensors(model: &Gguf, shape: &Shape) -> Result<(), String> { + match shape.family { + ModelFamily::DeepSeek => validate_deepseek_tensors(model, shape), + ModelFamily::Glm => validate_glm_tensors(model, shape), + } +} + +fn validate_deepseek_tensors(model: &Gguf, shape: &Shape) -> Result<(), String> { + let hc_dim = shape.embd * shape.hc; + let hc_mix = 2 * shape.hc + shape.hc * shape.hc; + let q_dim = shape.heads * shape.head_dim; + let output_low = shape.out_groups * shape.lora_o; + expect( + model, + "token_embd.weight", + &[F16], + &[shape.embd, shape.vocab], + )?; + expect(model, "output_hc_base.weight", &[F32], &[shape.hc])?; + expect(model, "output_hc_fn.weight", &[F16], &[hc_dim, shape.hc])?; + expect(model, "output_hc_scale.weight", &[F32], &[1])?; + expect(model, "output_norm.weight", &[F32], &[shape.embd])?; + expect(model, "output.weight", DENSE, &[shape.embd, shape.vocab])?; + + for layer in 0..shape.layers { + let name = |suffix: &str| format!("blk.{layer}.{suffix}"); + expect(model, &name("hc_attn_fn.weight"), &[F16], &[hc_dim, hc_mix])?; + expect(model, &name("hc_attn_scale.weight"), &[F32], &[3])?; + expect(model, &name("hc_attn_base.weight"), &[F32], &[hc_mix])?; + expect(model, &name("attn_norm.weight"), &[F32], &[shape.embd])?; + expect( + model, + &name("attn_q_a.weight"), + DENSE, + &[shape.embd, shape.lora_q], + )?; + expect( + model, + &name("attn_q_a_norm.weight"), + &[F32], + &[shape.lora_q], + )?; + expect( + model, + &name("attn_q_b.weight"), + DENSE, + &[shape.lora_q, q_dim], + )?; + expect( + model, + &name("attn_kv.weight"), + DENSE, + &[shape.embd, shape.head_dim], + )?; + expect( + model, + &name("attn_kv_a_norm.weight"), + &[F32], + &[shape.head_dim], + )?; + expect(model, &name("attn_sinks.weight"), &[F32], &[shape.heads])?; + expect( + model, + &name("attn_output_a.weight"), + DENSE, + &[ + shape.head_dim * (shape.heads / shape.out_groups), + output_low, + ], + )?; + expect( + model, + &name("attn_output_b.weight"), + DENSE, + &[output_low, shape.embd], + )?; + + let ratio = compression_ratio(shape, layer); + if ratio != 0 { + let compression_width = if ratio == 4 { 2 } else { 1 } * shape.head_dim; + expect( + model, + &name("attn_compressor_ape.weight"), + &[F16], + &[compression_width, u64::from(ratio)], + )?; + expect( + model, + &name("attn_compressor_kv.weight"), + &[F16], + &[shape.embd, compression_width], + )?; + expect( + model, + &name("attn_compressor_gate.weight"), + &[F16], + &[shape.embd, compression_width], + )?; + expect( + model, + &name("attn_compressor_norm.weight"), + &[F32], + &[shape.head_dim], + )?; + } + if ratio == 4 { + let index_q = shape.indexer_heads * shape.indexer_head_dim; + let index_width = 2 * shape.indexer_head_dim; + expect( + model, + &name("indexer.attn_q_b.weight"), + &[F16, Q8_0], + &[shape.lora_q, index_q], + )?; + expect( + model, + &name("indexer.proj.weight"), + &[F16], + &[shape.embd, shape.indexer_heads], + )?; + expect( + model, + &name("indexer_compressor_ape.weight"), + &[F16], + &[index_width, 4], + )?; + expect( + model, + &name("indexer_compressor_kv.weight"), + &[F16], + &[shape.embd, index_width], + )?; + expect( + model, + &name("indexer_compressor_gate.weight"), + &[F16], + &[shape.embd, index_width], + )?; + expect( + model, + &name("indexer_compressor_norm.weight"), + &[F32], + &[shape.indexer_head_dim], + )?; + } + expect(model, &name("hc_ffn_fn.weight"), &[F16], &[hc_dim, hc_mix])?; + expect(model, &name("hc_ffn_scale.weight"), &[F32], &[3])?; + expect(model, &name("hc_ffn_base.weight"), &[F32], &[hc_mix])?; + expect(model, &name("ffn_norm.weight"), &[F32], &[shape.embd])?; + expect( + model, + &name("ffn_gate_inp.weight"), + &[F16], + &[shape.embd, shape.experts], + )?; + expect_optional(model, &name("exp_probs_b.bias"), &[F32], &[shape.experts])?; + expect( + model, + &name("ffn_gate_exps.weight"), + ROUTED, + &[shape.embd, shape.ff_expert, shape.experts], + )?; + expect( + model, + &name("ffn_up_exps.weight"), + ROUTED, + &[shape.embd, shape.ff_expert, shape.experts], + )?; + expect( + model, + &name("ffn_down_exps.weight"), + ROUTED, + &[shape.ff_expert, shape.embd, shape.experts], + )?; + same_type( + model, + &name("ffn_gate_exps.weight"), + &name("ffn_up_exps.weight"), + )?; + expect( + model, + &name("ffn_gate_shexp.weight"), + DENSE, + &[shape.embd, shape.ff_expert], + )?; + expect( + model, + &name("ffn_up_shexp.weight"), + DENSE, + &[shape.embd, shape.ff_expert], + )?; + expect( + model, + &name("ffn_down_shexp.weight"), + DENSE, + &[shape.ff_expert, shape.embd], + )?; + if layer < shape.hash_layers { + expect( + model, + &name("ffn_gate_tid2eid.weight"), + &[I32], + &[shape.experts_used, shape.vocab], + )?; + } + } + Ok(()) +} + +fn validate_glm_tensors(model: &Gguf, shape: &Shape) -> Result<(), String> { + let q_dim = shape.heads * shape.key_mla; + let q_nope = shape.key_mla - shape.rot; + let index_q = shape.indexer_heads * shape.indexer_head_dim; + expect( + model, + "token_embd.weight", + DENSE, + &[shape.embd, shape.vocab], + )?; + expect(model, "output_norm.weight", &[F32], &[shape.embd])?; + expect(model, "output.weight", DENSE, &[shape.embd, shape.vocab])?; + for layer in 0..shape.layers { + let name = |suffix: &str| format!("blk.{layer}.{suffix}"); + expect(model, &name("attn_norm.weight"), &[F32], &[shape.embd])?; + expect( + model, + &name("attn_q_a.weight"), + DENSE, + &[shape.embd, shape.lora_q], + )?; + expect( + model, + &name("attn_q_a_norm.weight"), + &[F32], + &[shape.lora_q], + )?; + expect( + model, + &name("attn_q_b.weight"), + DENSE, + &[shape.lora_q, q_dim], + )?; + expect( + model, + &name("attn_kv_a_mqa.weight"), + DENSE, + &[shape.embd, shape.head_dim], + )?; + expect( + model, + &name("attn_kv_a_norm.weight"), + &[F32], + &[shape.kv_lora], + )?; + expect( + model, + &name("attn_k_b.weight"), + DENSE, + &[q_nope, shape.kv_lora, shape.heads], + )?; + expect( + model, + &name("attn_v_b.weight"), + DENSE, + &[shape.kv_lora, shape.value_mla, shape.heads], + )?; + expect( + model, + &name("attn_output.weight"), + DENSE, + &[shape.heads * shape.value_mla, shape.embd], + )?; + expect( + model, + &name("indexer.attn_k.weight"), + DENSE, + &[shape.embd, shape.indexer_head_dim], + )?; + expect( + model, + &name("indexer.attn_q_b.weight"), + DENSE, + &[shape.lora_q, index_q], + )?; + expect( + model, + &name("indexer.k_norm.weight"), + &[F32], + &[shape.indexer_head_dim], + )?; + expect( + model, + &name("indexer.k_norm.bias"), + &[F32], + &[shape.indexer_head_dim], + )?; + expect( + model, + &name("indexer.proj.weight"), + &[F32], + &[shape.embd, shape.indexer_heads], + )?; + expect(model, &name("ffn_norm.weight"), &[F32], &[shape.embd])?; + if layer < shape.leading_dense { + expect( + model, + &name("ffn_gate.weight"), + DENSE, + &[shape.embd, shape.ff_dense], + )?; + expect( + model, + &name("ffn_up.weight"), + DENSE, + &[shape.embd, shape.ff_dense], + )?; + expect( + model, + &name("ffn_down.weight"), + DENSE, + &[shape.ff_dense, shape.embd], + )?; + } else { + expect( + model, + &name("ffn_gate_inp.weight"), + &[F32], + &[shape.embd, shape.experts], + )?; + expect(model, &name("exp_probs_b.bias"), &[F32], &[shape.experts])?; + expect( + model, + &name("ffn_gate_exps.weight"), + ROUTED, + &[shape.embd, shape.ff_expert, shape.experts], + )?; + expect( + model, + &name("ffn_up_exps.weight"), + ROUTED, + &[shape.embd, shape.ff_expert, shape.experts], + )?; + expect( + model, + &name("ffn_down_exps.weight"), + ROUTED, + &[shape.ff_expert, shape.embd, shape.experts], + )?; + same_type( + model, + &name("ffn_gate_exps.weight"), + &name("ffn_up_exps.weight"), + )?; + expect( + model, + &name("ffn_gate_shexp.weight"), + DENSE, + &[shape.embd, shape.ff_expert], + )?; + expect( + model, + &name("ffn_up_shexp.weight"), + DENSE, + &[shape.embd, shape.ff_expert], + )?; + expect( + model, + &name("ffn_down_shexp.weight"), + DENSE, + &[shape.ff_expert, shape.embd], + )?; + } + if layer + shape.nextn >= shape.layers { + expect( + model, + &name("nextn.eh_proj.weight"), + DENSE, + &[2 * shape.embd, shape.embd], + )?; + expect(model, &name("nextn.enorm.weight"), &[F32], &[shape.embd])?; + expect(model, &name("nextn.hnorm.weight"), &[F32], &[shape.embd])?; + expect( + model, + &name("nextn.shared_head_norm.weight"), + &[F32], + &[shape.embd], + )?; + } + } + Ok(()) +} + +pub(super) fn validate_dspark(model: &Gguf, shape: &Shape) -> Result<(), String> { + if shape.model != ModelChoice::DeepSeekV4Flash { + return Err("DSpark support is available only for DeepSeek V4 Flash".into()); + } + let block_size = first_u32( + model, + &[ + "deepseek4.dspark.block_size", + "deepseek4.dspark_block_size", + "dspark.block_size", + ], + )?; + let markov_rank = first_u32( + model, + &[ + "deepseek4.dspark.markov_rank", + "deepseek4.dspark_markov_rank", + "dspark.markov_rank", + ], + )?; + let noise_token = first_u32( + model, + &[ + "deepseek4.dspark.noise_token_id", + "deepseek4.dspark_noise_token_id", + "dspark.noise_token_id", + ], + )?; + let targets = first_u32s( + model, + &[ + "deepseek4.dspark.target_layer_ids", + "deepseek4.dspark_target_layer_ids", + "dspark.target_layer_ids", + ], + )?; + if !(1..=16).contains(&block_size) || markov_rank == 0 || noise_token >= shape.vocab as u32 { + return Err("invalid DSpark block, Markov, or noise-token metadata".into()); + } + if targets.is_empty() + || targets.len() > 8 + || targets.windows(2).any(|pair| pair[0] >= pair[1]) + || targets.iter().any(|layer| *layer >= shape.layers) + { + return Err("invalid DSpark target-layer metadata".into()); + } + let stages = model + .tensors + .keys() + .filter_map(|name| { + name.strip_prefix("mtp.")? + .split('.') + .next()? + .parse::() + .ok() + }) + .max() + .map_or(0, |stage| stage + 1); + if !(1..=8).contains(&stages) { + return Err(format!("invalid DSpark stage count: {stages}")); + } + for stage in 0..stages { + validate_dspark_block(model, shape, stage)?; + if stage == 0 { + expect( + model, + &format!("mtp.{stage}.main_proj.weight"), + DSPARK_DENSE, + &[targets.len() as u64 * shape.embd, shape.embd], + )?; + expect( + model, + &format!("mtp.{stage}.main_norm.weight"), + &[F32], + &[shape.embd], + )?; + } + } + let prefix = format!("mtp.{}", stages - 1); + expect( + model, + &format!("{prefix}.norm.weight"), + &[F32], + &[shape.embd], + )?; + expect( + model, + &format!("{prefix}.hc_head_base.weight"), + &[F32], + &[shape.hc], + )?; + expect( + model, + &format!("{prefix}.hc_head_fn.weight"), + PLAIN, + &[shape.embd * shape.hc, shape.hc], + )?; + expect( + model, + &format!("{prefix}.hc_head_scale.weight"), + &[F32], + &[1], + )?; + expect( + model, + &format!("{prefix}.markov_head.markov_w1.weight"), + DSPARK_DENSE, + &[u64::from(markov_rank), shape.vocab], + )?; + expect( + model, + &format!("{prefix}.markov_head.markov_w2.weight"), + DSPARK_DENSE, + &[u64::from(markov_rank), shape.vocab], + )?; + expect( + model, + &format!("{prefix}.confidence_head.proj.weight"), + DSPARK_DENSE, + &[shape.embd + u64::from(markov_rank), 1], + )?; + Ok(()) +} + +fn validate_dspark_block(model: &Gguf, shape: &Shape, stage: u32) -> Result<(), String> { + let hc_dim = shape.embd * shape.hc; + let hc_mix = 2 * shape.hc + shape.hc * shape.hc; + let q_dim = shape.heads * shape.head_dim; + let output_low = shape.out_groups * shape.lora_o; + let name = |suffix: &str| format!("mtp.{stage}.{suffix}"); + for (suffix, types, dims) in [ + ("hc_attn_fn.weight", PLAIN, vec![hc_dim, hc_mix]), + ("hc_attn_scale.weight", &[F32][..], vec![3]), + ("hc_attn_base.weight", &[F32][..], vec![hc_mix]), + ("attn_norm.weight", &[F32][..], vec![shape.embd]), + ( + "attn_q_a.weight", + DSPARK_DENSE, + vec![shape.embd, shape.lora_q], + ), + ("attn_q_a_norm.weight", &[F32][..], vec![shape.lora_q]), + ("attn_q_b.weight", DSPARK_DENSE, vec![shape.lora_q, q_dim]), + ( + "attn_kv.weight", + DSPARK_DENSE, + vec![shape.embd, shape.head_dim], + ), + ("attn_kv_a_norm.weight", &[F32][..], vec![shape.head_dim]), + ("attn_sinks.weight", &[F32][..], vec![shape.heads]), + ( + "attn_output_a.weight", + DSPARK_DENSE, + vec![ + shape.head_dim * (shape.heads / shape.out_groups), + output_low, + ], + ), + ( + "attn_output_b.weight", + DSPARK_DENSE, + vec![output_low, shape.embd], + ), + ("hc_ffn_fn.weight", PLAIN, vec![hc_dim, hc_mix]), + ("hc_ffn_scale.weight", &[F32][..], vec![3]), + ("hc_ffn_base.weight", &[F32][..], vec![hc_mix]), + ("ffn_norm.weight", &[F32][..], vec![shape.embd]), + ( + "ffn_gate_inp.weight", + DSPARK_DENSE, + vec![shape.embd, shape.experts], + ), + ("exp_probs_b.bias", &[F32][..], vec![shape.experts]), + ( + "ffn_gate_exps.weight", + ROUTED, + vec![shape.embd, shape.ff_expert, shape.experts], + ), + ( + "ffn_up_exps.weight", + ROUTED, + vec![shape.embd, shape.ff_expert, shape.experts], + ), + ( + "ffn_down_exps.weight", + ROUTED, + vec![shape.ff_expert, shape.embd, shape.experts], + ), + ( + "ffn_gate_shexp.weight", + DSPARK_DENSE, + vec![shape.embd, shape.ff_expert], + ), + ( + "ffn_up_shexp.weight", + DSPARK_DENSE, + vec![shape.embd, shape.ff_expert], + ), + ( + "ffn_down_shexp.weight", + DSPARK_DENSE, + vec![shape.ff_expert, shape.embd], + ), + ] { + expect(model, &name(suffix), types, &dims)?; + } + same_type( + model, + &name("ffn_gate_exps.weight"), + &name("ffn_up_exps.weight"), + ) +} + +fn expect(model: &Gguf, name: &str, types: &[u32], dims: &[u64]) -> Result<(), String> { + validate_tensor(name, model.tensor(name)?, types, dims) +} + +fn expect_optional(model: &Gguf, name: &str, types: &[u32], dims: &[u64]) -> Result<(), String> { + match model.tensors.get(name) { + Some(tensor) => validate_tensor(name, tensor, types, dims), + None => Ok(()), + } +} + +fn validate_tensor(name: &str, tensor: &Tensor, types: &[u32], dims: &[u64]) -> Result<(), String> { + if !types.contains(&tensor.kind) { + return Err(format!( + "tensor {name} has unsupported type {}", + tensor.kind + )); + } + if tensor.dims != dims { + return Err(format!( + "tensor {name} has dimensions {:?}, expected {dims:?}", + tensor.dims + )); + } + Ok(()) +} + +fn same_type(model: &Gguf, first: &str, second: &str) -> Result<(), String> { + if model.tensor(first)?.kind != model.tensor(second)?.kind { + Err(format!( + "tensors {first} and {second} use different quantizations" + )) + } else { + Ok(()) + } +} + +fn expect_u64(model: &Gguf, key: &str, expected: u64) -> Result<(), String> { + let actual = model.u64(key)?; + if actual == expected { + Ok(()) + } else { + Err(format!("{key} is {actual}, expected {expected}")) + } +} + +fn expect_float(model: &Gguf, key: &str, expected: f32) -> Result<(), String> { + let actual = model.f32(key)?; + if float_eq(actual, expected) { + Ok(()) + } else { + Err(format!("{key} is {actual}, expected {expected}")) + } +} + +fn float_eq(actual: f32, expected: f32) -> bool { + actual.is_finite() && (actual - expected).abs() <= expected.abs().max(1.0) * 1.0e-6 +} + +fn compression_ratio(shape: &Shape, layer: u32) -> u32 { + match shape.model { + ModelChoice::DeepSeekV4Flash if layer < 2 => 0, + ModelChoice::DeepSeekV4Pro if layer < 2 => 128, + ModelChoice::DeepSeekV4Flash | ModelChoice::DeepSeekV4Pro if layer.is_multiple_of(2) => 4, + ModelChoice::DeepSeekV4Flash | ModelChoice::DeepSeekV4Pro => 128, + ModelChoice::Glm52 => 0, + } +} + +fn first_u32(model: &Gguf, keys: &[&str]) -> Result { + keys.iter() + .find_map(|key| model.u32(key).ok()) + .ok_or_else(|| format!("required DSpark metadata is missing: {}", keys[0])) +} + +fn first_u32s<'a>(model: &'a Gguf, keys: &[&str]) -> Result<&'a [u32], String> { + keys.iter() + .find_map(|key| model.u32s(key).ok()) + .ok_or_else(|| format!("required DSpark metadata is missing: {}", keys[0])) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn installed_ds4_fixture_opens_and_renders_a_prompt() { + let path = Path::new( + "../ds4/gguf/DeepSeek-V4-Flash-IQ2XXS-w2Q2K-AProjQ8-SExpQ8-OutQ8-chat-v2-imatrix.gguf", + ); + if !path.exists() { + return; + } + let model = Model::open_main(path, ModelChoice::DeepSeekV4Flash).unwrap(); + let summary = model.summary(); + assert_eq!(summary.model, ModelChoice::DeepSeekV4Flash); + assert_eq!(summary.vocabulary_size, 129_280); + assert_eq!( + model.tokenize("Hello, world! 1234\nint café = 7;\n中文テスト"), + [ + 19_923, 14, 2_058, 3, 223, 6_895, 22, 201, 650, 57_664, 438, 223, 25, 510, 21_134, + 109_288, + ] + ); + assert_eq!( + model.render_prompt("", "Hello", ReasoningMode::Direct), + [0, 128_803, 19_923, 128_804, 128_822] + ); + assert_eq!(model.tokenizer.tokenize_rendered("|DSML|").len(), 1); + assert!(model.is_stop_token_for_reasoning(128_822, ReasoningMode::Direct)); + assert!(!model.is_stop_token_for_reasoning(128_822, ReasoningMode::High)); + let think_start = *model + .render_prompt("", "Hello", ReasoningMode::High) + .last() + .unwrap(); + assert_eq!( + model.render_conversation( + "", + &[ + ChatTurn { + user: true, + skip_previous_eos: false, + reasoning: None, + reasoning_complete: true, + content: "Hello".into(), + }, + ChatTurn { + user: false, + skip_previous_eos: false, + reasoning: None, + reasoning_complete: true, + content: "Hello".into(), + }, + ChatTurn { + user: true, + skip_previous_eos: false, + reasoning: None, + reasoning_complete: true, + content: "Hello".into(), + }, + ], + ReasoningMode::Direct, + ), + [ + 0, 128_803, 19_923, 128_804, 128_822, 19_923, 1, 128_803, 19_923, 128_804, 128_822, + ] + ); + assert_eq!( + model.render_continuation("Hello", ReasoningMode::Direct, false), + [1, 128_803, 19_923, 128_804, 128_822] + ); + assert_eq!( + model.render_continuation("Hello", ReasoningMode::Direct, true), + [128_803, 19_923, 128_804, 128_822] + ); + assert_eq!( + model.render_conversation( + "", + &[ + ChatTurn { + user: true, + skip_previous_eos: false, + reasoning: None, + reasoning_complete: true, + content: "Hello".into(), + }, + ChatTurn { + user: false, + skip_previous_eos: false, + reasoning: Some("Hello".into()), + reasoning_complete: true, + content: "Hello".into(), + }, + ChatTurn { + user: true, + skip_previous_eos: false, + reasoning: None, + reasoning_complete: true, + content: "Hello".into(), + }, + ], + ReasoningMode::Direct, + ), + [ + 0, + 128_803, + 19_923, + 128_804, + think_start, + 19_923, + 128_822, + 19_923, + 1, + 128_803, + 19_923, + 128_804, + 128_822, + ] + ); + } + + #[test] + fn installed_dspark_fixture_passes_the_target_layout() { + let path = Path::new("../ds4/gguf/DeepSeek-V4-Flash-DSpark-support.gguf"); + if path.exists() { + validate_model_artifact(path, ModelChoice::DeepSeekV4Flash, true).unwrap(); + } + } +} diff --git a/src/model.rs b/src/model.rs index 8264e16..7f0c698 100644 --- a/src/model.rs +++ b/src/model.rs @@ -1,10 +1,12 @@ -use std::fmt; -use std::fs::{self, File, OpenOptions}; -use std::io::{Read, Write}; -use std::path::{Path, PathBuf}; -use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; +mod transfer; -use sha2::{Digest, Sha256}; +pub(crate) use transfer::{ + delete_managed_artifact, download_managed_artifact, validate_managed_artifact, +}; + +use std::fmt; +use std::fs; +use std::path::{Path, PathBuf}; pub(crate) const MODEL_CHOICES: [ModelChoice; 3] = [ ModelChoice::DeepSeekV4Flash, @@ -379,505 +381,3 @@ pub(crate) fn artifact_verification_progress( }), } } - -pub(crate) fn download_managed_artifact( - id: ManagedArtifactId, - models_path: &Path, - cancel: &AtomicBool, - verified_bytes: &AtomicU64, -) -> Result { - download_artifact_with_cancel( - id.model(), - id.artifact(), - models_path, - cancel, - verified_bytes, - ) -} - -pub(crate) fn validate_managed_artifact( - id: ManagedArtifactId, - models_path: &Path, - cancel: &AtomicBool, - verified_bytes: &AtomicU64, -) -> Result { - let model = id.model(); - let artifact = id.artifact(); - let destination = artifact.path(model, models_path); - let partial = artifact.partial_path(model, models_path); - let (path, promote) = if destination.exists() { - (destination.clone(), false) - } else if partial.exists() { - (partial.clone(), true) - } else { - return Err(format!("{} is not downloaded", artifact.label)); - }; - - match verify(&path, artifact, model, cancel, verified_bytes) { - Ok(DownloadOutcome::Stopped) => return Ok(DownloadOutcome::Stopped), - Ok(DownloadOutcome::Complete) => {} - Err(error) => { - let marker = artifact.verification_path(model, models_path); - if let Err(remove_error) = fs::remove_file(marker) - && remove_error.kind() != std::io::ErrorKind::NotFound - { - return Err(format!( - "{error}; could not remove checksum marker: {remove_error}" - )); - } - return Err(error); - } - } - if promote { - fs::rename(partial, destination).map_err(|error| error.to_string())?; - } - mark_verified(model, artifact, models_path)?; - Ok(DownloadOutcome::Complete) -} - -pub(crate) fn delete_managed_artifact( - id: ManagedArtifactId, - models_path: &Path, -) -> Result<(), String> { - let model = id.model(); - let artifact = id.artifact(); - for path in [ - artifact.path(model, models_path), - artifact.partial_path(model, models_path), - artifact.verification_path(model, models_path), - ] { - match fs::remove_file(path) { - Ok(()) => {} - Err(error) if error.kind() == std::io::ErrorKind::NotFound => {} - Err(error) => return Err(error.to_string()), - } - } - Ok(()) -} - -#[cfg(test)] -fn download_artifact( - model: ModelChoice, - artifact: &Artifact, - models_path: &Path, -) -> Result { - download_artifact_with_cancel( - model, - artifact, - models_path, - &AtomicBool::new(false), - &AtomicU64::new(0), - ) -} - -fn download_artifact_with_cancel( - model: ModelChoice, - artifact: &Artifact, - models_path: &Path, - cancel: &AtomicBool, - verified_bytes: &AtomicU64, -) -> Result { - if cancel.load(Ordering::Relaxed) { - return Ok(DownloadOutcome::Stopped); - } - let destination = artifact.path(model, models_path); - if artifact.is_installed(model, models_path) { - return Ok(DownloadOutcome::Complete); - } - if destination.exists() { - if verify(&destination, artifact, model, cancel, verified_bytes)? - == DownloadOutcome::Stopped - { - return Ok(DownloadOutcome::Stopped); - } - mark_verified(model, artifact, models_path)?; - return Ok(DownloadOutcome::Complete); - } - - let directory = destination - .parent() - .ok_or_else(|| "model artifact path has no parent directory".to_owned())?; - fs::create_dir_all(directory).map_err(|error| error.to_string())?; - let partial = artifact.partial_path(model, models_path); - let partial_size = partial.metadata().map_or(0, |metadata| metadata.len()); - if partial_size > artifact.size { - File::create(&partial).map_err(|error| error.to_string())?; - } - if partial.metadata().map_or(0, |metadata| metadata.len()) != artifact.size - && download_to_partial(artifact, &partial, cancel)? == DownloadOutcome::Stopped - { - return Ok(DownloadOutcome::Stopped); - } - - if cancel.load(Ordering::Relaxed) { - return Ok(DownloadOutcome::Stopped); - } - match verify(&partial, artifact, model, cancel, verified_bytes) { - Ok(DownloadOutcome::Stopped) => return Ok(DownloadOutcome::Stopped), - Ok(DownloadOutcome::Complete) => {} - Err(error) => { - fs::remove_file(&partial).map_err(|remove_error| { - format!("{error}; could not remove partial file: {remove_error}") - })?; - return Err(error); - } - } - fs::rename(partial, destination).map_err(|error| error.to_string())?; - mark_verified(model, artifact, models_path)?; - Ok(DownloadOutcome::Complete) -} - -fn mark_verified( - model: ModelChoice, - artifact: &Artifact, - models_path: &Path, -) -> Result<(), String> { - fs::write( - artifact.verification_path(model, models_path), - artifact.sha256, - ) - .map_err(|error| error.to_string()) -} - -fn verify( - path: &Path, - artifact: &Artifact, - model: ModelChoice, - cancel: &AtomicBool, - verified_bytes: &AtomicU64, -) -> Result { - verified_bytes.store(0, Ordering::Relaxed); - let size = path.metadata().map_err(|error| error.to_string())?.len(); - if size != artifact.size { - return Err(format!( - "{} has size {size}, expected {}", - path.display(), - artifact.size - )); - } - - let mut file = File::open(path).map_err(|error| error.to_string())?; - let mut hasher = Sha256::new(); - let mut buffer = vec![0; 1024 * 1024]; - loop { - if cancel.load(Ordering::Relaxed) { - return Ok(DownloadOutcome::Stopped); - } - let count = file.read(&mut buffer).map_err(|error| error.to_string())?; - if count == 0 { - break; - } - hasher.update(&buffer[..count]); - verified_bytes.fetch_add(count as u64, Ordering::Relaxed); - } - let actual = hex(&hasher.finalize()); - if actual != artifact.sha256 { - return Err(format!( - "Checksum verification failed for {}", - path.display() - )); - } - if let Some(support) = artifact.support { - crate::engine::validate_model_artifact(path, model, support)?; - } - Ok(DownloadOutcome::Complete) -} - -fn hex(bytes: &[u8]) -> String { - const DIGITS: &[u8; 16] = b"0123456789abcdef"; - let mut result = String::with_capacity(bytes.len() * 2); - for byte in bytes { - result.push(DIGITS[(byte >> 4) as usize] as char); - result.push(DIGITS[(byte & 0x0f) as usize] as char); - } - result -} - -fn download_to_partial( - artifact: &Artifact, - partial: &Path, - cancel: &AtomicBool, -) -> Result { - download_url_to_partial(&artifact.url(), partial, cancel) -} - -fn download_url_to_partial( - url: &str, - partial: &Path, - cancel: &AtomicBool, -) -> Result { - let offset = partial.metadata().map_or(0, |metadata| metadata.len()); - let agent: ureq::Agent = ureq::Agent::config_builder() - .https_only(url.starts_with("https://")) - .build() - .into(); - let mut request = agent.get(url); - if offset > 0 { - request = request.header("Range", format!("bytes={offset}-")); - } - let mut response = request - .call() - .map_err(|error| format!("Model download failed: {error}"))?; - let status = response.status().as_u16(); - let append = offset > 0 && status == 206; - if offset > 0 && status != 200 && status != 206 { - return Err(format!( - "Model server returned HTTP {status} while resuming at byte {offset}" - )); - } - if append { - validate_content_range(&response, offset)?; - } - - let mut output = OpenOptions::new() - .create(true) - .write(true) - .append(append) - .truncate(!append) - .open(partial) - .map_err(|error| error.to_string())?; - let mut body = response.body_mut().as_reader(); - let mut buffer = vec![0; 1024 * 1024]; - loop { - if cancel.load(Ordering::Relaxed) { - return Ok(DownloadOutcome::Stopped); - } - let count = body - .read(&mut buffer) - .map_err(|error| format!("Model download failed: {error}"))?; - if count == 0 { - break; - } - output - .write_all(&buffer[..count]) - .map_err(|error| error.to_string())?; - } - Ok(DownloadOutcome::Complete) -} - -fn validate_content_range( - response: &ureq::http::Response, - offset: u64, -) -> Result<(), String> { - let expected = format!("bytes {offset}-"); - let content_range = response - .headers() - .get("content-range") - .and_then(|value| value.to_str().ok()) - .unwrap_or_default(); - if content_range.starts_with(&expected) { - Ok(()) - } else { - Err(format!( - "Model server returned an invalid Content-Range while resuming at byte {offset}" - )) - } -} - -#[cfg(test)] -mod tests { - use super::*; - use std::net::TcpListener; - use std::thread; - use std::time::{SystemTime, UNIX_EPOCH}; - - #[test] - fn catalog_and_checksum_verification_are_explicit() { - assert_eq!(ModelChoice::from_id("glm-5.2"), Some(ModelChoice::Glm52)); - assert!(ModelChoice::from_id("unknown").is_none()); - assert_eq!( - ModelChoice::DeepSeekV4Flash.main_artifact().size, - 86_720_111_488 - ); - assert_eq!(ModelChoice::Glm52.main_artifact().size, 211_075_856_448); - assert_eq!(ModelChoice::DeepSeekV4Flash.artifacts(true).count(), 2); - assert_eq!(ModelChoice::Glm52.artifacts(true).count(), 1); - - let id = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_nanos(); - let models_path = std::env::temp_dir().join(format!("ds4-server-models-{id}")); - let engine = engine_artifacts(ModelChoice::DeepSeekV4Flash, true, &models_path); - assert_eq!( - engine.model.file_name(), - Some(std::ffi::OsStr::new(FLASH.file_name)) - ); - assert_eq!( - engine.mtp.as_deref().and_then(Path::file_name), - Some(std::ffi::OsStr::new(FLASH_DSPARK.file_name)) - ); - let empty = Artifact { - label: "empty", - file_name: "empty", - repository: "", - size: 0, - sha256: "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855", - support: None, - }; - let partial = empty.partial_path(ModelChoice::DeepSeekV4Flash, &models_path); - fs::create_dir_all(partial.parent().unwrap()).unwrap(); - fs::write(&partial, []).unwrap(); - download_artifact(ModelChoice::DeepSeekV4Flash, &empty, &models_path).unwrap(); - assert!(empty.is_installed(ModelChoice::DeepSeekV4Flash, &models_path)); - assert!(!partial.exists()); - assert_eq!( - fs::read_to_string(empty.verification_path(ModelChoice::DeepSeekV4Flash, &models_path)) - .unwrap(), - empty.sha256 - ); - fs::remove_dir_all(models_path).unwrap(); - } - - #[test] - fn verification_reports_bytes_read() { - let path = std::env::temp_dir().join(format!( - "ds4-server-verify-{}", - SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_nanos() - )); - fs::write(&path, b"abc").unwrap(); - let artifact = Artifact { - label: "test model", - file_name: "unused", - repository: "unused", - size: 3, - sha256: "ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad", - support: None, - }; - let verified_bytes = AtomicU64::new(999); - - assert_eq!( - verify( - &path, - &artifact, - ModelChoice::DeepSeekV4Flash, - &AtomicBool::new(false), - &verified_bytes, - ) - .unwrap(), - DownloadOutcome::Complete - ); - assert_eq!(verified_bytes.load(Ordering::Relaxed), 3); - fs::remove_file(path).unwrap(); - } - - #[test] - fn managed_artifact_inventory_and_delete_include_partial_files() { - let models_path = std::env::temp_dir().join(format!( - "ds4-server-manager-{}", - SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_nanos() - )); - let id = ManagedArtifactId::DeepSeekV4Flash; - let partial = id.artifact().partial_path(id.model(), &models_path); - fs::create_dir_all(partial.parent().unwrap()).unwrap(); - fs::write(&partial, b"part").unwrap(); - - let managed = managed_artifacts(&models_path) - .into_iter() - .find(|artifact| artifact.id == id) - .unwrap(); - assert_eq!(managed.stored, 4); - assert_eq!(managed.state, ManagedArtifactState::Partial); - - delete_managed_artifact(id, &models_path).unwrap(); - assert!(!partial.exists()); - fs::remove_dir_all(models_path).unwrap(); - } - - #[test] - fn restart_resumes_at_the_existing_partial_byte() { - let content = b"restart-resume works"; - let offset = 8; - let directory = std::env::temp_dir().join(format!( - "ds4-server-resume-{}", - SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_nanos() - )); - fs::create_dir_all(&directory).unwrap(); - let partial = directory.join("model.gguf.part"); - fs::write(&partial, &content[..offset]).unwrap(); - - let listener = match TcpListener::bind("127.0.0.1:0") { - Ok(listener) => listener, - Err(error) if error.kind() == std::io::ErrorKind::PermissionDenied => { - fs::remove_dir_all(directory).unwrap(); - return; - } - Err(error) => panic!("could not start test server: {error}"), - }; - let address = listener.local_addr().unwrap(); - let server = thread::spawn(move || { - let (mut connection, _) = listener.accept().unwrap(); - let mut request = [0; 2048]; - let count = connection.read(&mut request).unwrap(); - let request = String::from_utf8_lossy(&request[..count]).to_ascii_lowercase(); - assert!(request.contains("range: bytes=8-")); - let remaining = &content[offset..]; - write!( - connection, - "HTTP/1.1 206 Partial Content\r\nContent-Length: {}\r\nContent-Range: bytes {offset}-{}/{}\r\nConnection: close\r\n\r\n", - remaining.len(), - content.len() - 1, - content.len(), - ) - .unwrap(); - connection.write_all(remaining).unwrap(); - }); - - let outcome = download_url_to_partial( - &format!("http://{address}/model.gguf"), - &partial, - &AtomicBool::new(false), - ) - .unwrap(); - server.join().unwrap(); - assert_eq!(outcome, DownloadOutcome::Complete); - assert_eq!(fs::read(&partial).unwrap(), content); - fs::remove_dir_all(directory).unwrap(); - } - - #[test] - fn cancellation_keeps_the_partial_file_for_the_next_run() { - let directory = std::env::temp_dir().join(format!( - "ds4-server-cancel-{}", - SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap() - .as_nanos() - )); - let artifact = Artifact { - label: "test model", - file_name: "model.gguf", - repository: "unused", - size: 10, - sha256: "unused", - support: None, - }; - let partial = artifact.partial_path(ModelChoice::DeepSeekV4Flash, &directory); - fs::create_dir_all(partial.parent().unwrap()).unwrap(); - fs::write(&partial, b"part").unwrap(); - let cancel = AtomicBool::new(true); - - assert_eq!( - download_artifact_with_cancel( - ModelChoice::DeepSeekV4Flash, - &artifact, - &directory, - &cancel, - &AtomicU64::new(0), - ) - .unwrap(), - DownloadOutcome::Stopped - ); - assert_eq!(fs::read(&partial).unwrap(), b"part"); - fs::remove_dir_all(directory).unwrap(); - } -} diff --git a/src/model/transfer.rs b/src/model/transfer.rs new file mode 100644 index 0000000..2854e6a --- /dev/null +++ b/src/model/transfer.rs @@ -0,0 +1,507 @@ +use super::*; +use sha2::{Digest, Sha256}; +use std::fs::{self, File, OpenOptions}; +use std::io::{Read, Write}; +use std::sync::atomic::{AtomicBool, AtomicU64, Ordering}; + +pub(crate) fn download_managed_artifact( + id: ManagedArtifactId, + models_path: &Path, + cancel: &AtomicBool, + verified_bytes: &AtomicU64, +) -> Result { + download_artifact_with_cancel( + id.model(), + id.artifact(), + models_path, + cancel, + verified_bytes, + ) +} + +pub(crate) fn validate_managed_artifact( + id: ManagedArtifactId, + models_path: &Path, + cancel: &AtomicBool, + verified_bytes: &AtomicU64, +) -> Result { + let model = id.model(); + let artifact = id.artifact(); + let destination = artifact.path(model, models_path); + let partial = artifact.partial_path(model, models_path); + let (path, promote) = if destination.exists() { + (destination.clone(), false) + } else if partial.exists() { + (partial.clone(), true) + } else { + return Err(format!("{} is not downloaded", artifact.label)); + }; + + match verify(&path, artifact, model, cancel, verified_bytes) { + Ok(DownloadOutcome::Stopped) => return Ok(DownloadOutcome::Stopped), + Ok(DownloadOutcome::Complete) => {} + Err(error) => { + let marker = artifact.verification_path(model, models_path); + if let Err(remove_error) = fs::remove_file(marker) + && remove_error.kind() != std::io::ErrorKind::NotFound + { + return Err(format!( + "{error}; could not remove checksum marker: {remove_error}" + )); + } + return Err(error); + } + } + if promote { + fs::rename(partial, destination).map_err(|error| error.to_string())?; + } + mark_verified(model, artifact, models_path)?; + Ok(DownloadOutcome::Complete) +} + +pub(crate) fn delete_managed_artifact( + id: ManagedArtifactId, + models_path: &Path, +) -> Result<(), String> { + let model = id.model(); + let artifact = id.artifact(); + for path in [ + artifact.path(model, models_path), + artifact.partial_path(model, models_path), + artifact.verification_path(model, models_path), + ] { + match fs::remove_file(path) { + Ok(()) => {} + Err(error) if error.kind() == std::io::ErrorKind::NotFound => {} + Err(error) => return Err(error.to_string()), + } + } + Ok(()) +} + +#[cfg(test)] +fn download_artifact( + model: ModelChoice, + artifact: &Artifact, + models_path: &Path, +) -> Result { + download_artifact_with_cancel( + model, + artifact, + models_path, + &AtomicBool::new(false), + &AtomicU64::new(0), + ) +} + +fn download_artifact_with_cancel( + model: ModelChoice, + artifact: &Artifact, + models_path: &Path, + cancel: &AtomicBool, + verified_bytes: &AtomicU64, +) -> Result { + if cancel.load(Ordering::Relaxed) { + return Ok(DownloadOutcome::Stopped); + } + let destination = artifact.path(model, models_path); + if artifact.is_installed(model, models_path) { + return Ok(DownloadOutcome::Complete); + } + if destination.exists() { + if verify(&destination, artifact, model, cancel, verified_bytes)? + == DownloadOutcome::Stopped + { + return Ok(DownloadOutcome::Stopped); + } + mark_verified(model, artifact, models_path)?; + return Ok(DownloadOutcome::Complete); + } + + let directory = destination + .parent() + .ok_or_else(|| "model artifact path has no parent directory".to_owned())?; + fs::create_dir_all(directory).map_err(|error| error.to_string())?; + let partial = artifact.partial_path(model, models_path); + let partial_size = partial.metadata().map_or(0, |metadata| metadata.len()); + if partial_size > artifact.size { + File::create(&partial).map_err(|error| error.to_string())?; + } + if partial.metadata().map_or(0, |metadata| metadata.len()) != artifact.size + && download_to_partial(artifact, &partial, cancel)? == DownloadOutcome::Stopped + { + return Ok(DownloadOutcome::Stopped); + } + + if cancel.load(Ordering::Relaxed) { + return Ok(DownloadOutcome::Stopped); + } + match verify(&partial, artifact, model, cancel, verified_bytes) { + Ok(DownloadOutcome::Stopped) => return Ok(DownloadOutcome::Stopped), + Ok(DownloadOutcome::Complete) => {} + Err(error) => { + fs::remove_file(&partial).map_err(|remove_error| { + format!("{error}; could not remove partial file: {remove_error}") + })?; + return Err(error); + } + } + fs::rename(partial, destination).map_err(|error| error.to_string())?; + mark_verified(model, artifact, models_path)?; + Ok(DownloadOutcome::Complete) +} + +fn mark_verified( + model: ModelChoice, + artifact: &Artifact, + models_path: &Path, +) -> Result<(), String> { + fs::write( + artifact.verification_path(model, models_path), + artifact.sha256, + ) + .map_err(|error| error.to_string()) +} + +fn verify( + path: &Path, + artifact: &Artifact, + model: ModelChoice, + cancel: &AtomicBool, + verified_bytes: &AtomicU64, +) -> Result { + verified_bytes.store(0, Ordering::Relaxed); + let size = path.metadata().map_err(|error| error.to_string())?.len(); + if size != artifact.size { + return Err(format!( + "{} has size {size}, expected {}", + path.display(), + artifact.size + )); + } + + let mut file = File::open(path).map_err(|error| error.to_string())?; + let mut hasher = Sha256::new(); + let mut buffer = vec![0; 1024 * 1024]; + loop { + if cancel.load(Ordering::Relaxed) { + return Ok(DownloadOutcome::Stopped); + } + let count = file.read(&mut buffer).map_err(|error| error.to_string())?; + if count == 0 { + break; + } + hasher.update(&buffer[..count]); + verified_bytes.fetch_add(count as u64, Ordering::Relaxed); + } + let actual = hex(&hasher.finalize()); + if actual != artifact.sha256 { + return Err(format!( + "Checksum verification failed for {}", + path.display() + )); + } + if let Some(support) = artifact.support { + crate::engine::validate_model_artifact(path, model, support)?; + } + Ok(DownloadOutcome::Complete) +} + +fn hex(bytes: &[u8]) -> String { + const DIGITS: &[u8; 16] = b"0123456789abcdef"; + let mut result = String::with_capacity(bytes.len() * 2); + for byte in bytes { + result.push(DIGITS[(byte >> 4) as usize] as char); + result.push(DIGITS[(byte & 0x0f) as usize] as char); + } + result +} + +fn download_to_partial( + artifact: &Artifact, + partial: &Path, + cancel: &AtomicBool, +) -> Result { + download_url_to_partial(&artifact.url(), partial, cancel) +} + +fn download_url_to_partial( + url: &str, + partial: &Path, + cancel: &AtomicBool, +) -> Result { + let offset = partial.metadata().map_or(0, |metadata| metadata.len()); + let agent: ureq::Agent = ureq::Agent::config_builder() + .https_only(url.starts_with("https://")) + .build() + .into(); + let mut request = agent.get(url); + if offset > 0 { + request = request.header("Range", format!("bytes={offset}-")); + } + let mut response = request + .call() + .map_err(|error| format!("Model download failed: {error}"))?; + let status = response.status().as_u16(); + let append = offset > 0 && status == 206; + if offset > 0 && status != 200 && status != 206 { + return Err(format!( + "Model server returned HTTP {status} while resuming at byte {offset}" + )); + } + if append { + validate_content_range(&response, offset)?; + } + + let mut output = OpenOptions::new() + .create(true) + .write(true) + .append(append) + .truncate(!append) + .open(partial) + .map_err(|error| error.to_string())?; + let mut body = response.body_mut().as_reader(); + let mut buffer = vec![0; 1024 * 1024]; + loop { + if cancel.load(Ordering::Relaxed) { + return Ok(DownloadOutcome::Stopped); + } + let count = body + .read(&mut buffer) + .map_err(|error| format!("Model download failed: {error}"))?; + if count == 0 { + break; + } + output + .write_all(&buffer[..count]) + .map_err(|error| error.to_string())?; + } + Ok(DownloadOutcome::Complete) +} + +fn validate_content_range( + response: &ureq::http::Response, + offset: u64, +) -> Result<(), String> { + let expected = format!("bytes {offset}-"); + let content_range = response + .headers() + .get("content-range") + .and_then(|value| value.to_str().ok()) + .unwrap_or_default(); + if content_range.starts_with(&expected) { + Ok(()) + } else { + Err(format!( + "Model server returned an invalid Content-Range while resuming at byte {offset}" + )) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::net::TcpListener; + use std::thread; + use std::time::{SystemTime, UNIX_EPOCH}; + + #[test] + fn catalog_and_checksum_verification_are_explicit() { + assert_eq!(ModelChoice::from_id("glm-5.2"), Some(ModelChoice::Glm52)); + assert!(ModelChoice::from_id("unknown").is_none()); + assert_eq!( + ModelChoice::DeepSeekV4Flash.main_artifact().size, + 86_720_111_488 + ); + assert_eq!(ModelChoice::Glm52.main_artifact().size, 211_075_856_448); + assert_eq!(ModelChoice::DeepSeekV4Flash.artifacts(true).count(), 2); + assert_eq!(ModelChoice::Glm52.artifacts(true).count(), 1); + + let id = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_nanos(); + let models_path = std::env::temp_dir().join(format!("ds4-server-models-{id}")); + let engine = engine_artifacts(ModelChoice::DeepSeekV4Flash, true, &models_path); + assert_eq!( + engine.model.file_name(), + Some(std::ffi::OsStr::new(FLASH.file_name)) + ); + assert_eq!( + engine.mtp.as_deref().and_then(Path::file_name), + Some(std::ffi::OsStr::new(FLASH_DSPARK.file_name)) + ); + let empty = Artifact { + label: "empty", + file_name: "empty", + repository: "", + size: 0, + sha256: "e3b0c44298fc1c149afbf4c8996fb92427ae41e4649b934ca495991b7852b855", + support: None, + }; + let partial = empty.partial_path(ModelChoice::DeepSeekV4Flash, &models_path); + fs::create_dir_all(partial.parent().unwrap()).unwrap(); + fs::write(&partial, []).unwrap(); + download_artifact(ModelChoice::DeepSeekV4Flash, &empty, &models_path).unwrap(); + assert!(empty.is_installed(ModelChoice::DeepSeekV4Flash, &models_path)); + assert!(!partial.exists()); + assert_eq!( + fs::read_to_string(empty.verification_path(ModelChoice::DeepSeekV4Flash, &models_path)) + .unwrap(), + empty.sha256 + ); + fs::remove_dir_all(models_path).unwrap(); + } + + #[test] + fn verification_reports_bytes_read() { + let path = std::env::temp_dir().join(format!( + "ds4-server-verify-{}", + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_nanos() + )); + fs::write(&path, b"abc").unwrap(); + let artifact = Artifact { + label: "test model", + file_name: "unused", + repository: "unused", + size: 3, + sha256: "ba7816bf8f01cfea414140de5dae2223b00361a396177a9cb410ff61f20015ad", + support: None, + }; + let verified_bytes = AtomicU64::new(999); + + assert_eq!( + verify( + &path, + &artifact, + ModelChoice::DeepSeekV4Flash, + &AtomicBool::new(false), + &verified_bytes, + ) + .unwrap(), + DownloadOutcome::Complete + ); + assert_eq!(verified_bytes.load(Ordering::Relaxed), 3); + fs::remove_file(path).unwrap(); + } + + #[test] + fn managed_artifact_inventory_and_delete_include_partial_files() { + let models_path = std::env::temp_dir().join(format!( + "ds4-server-manager-{}", + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_nanos() + )); + let id = ManagedArtifactId::DeepSeekV4Flash; + let partial = id.artifact().partial_path(id.model(), &models_path); + fs::create_dir_all(partial.parent().unwrap()).unwrap(); + fs::write(&partial, b"part").unwrap(); + + let managed = managed_artifacts(&models_path) + .into_iter() + .find(|artifact| artifact.id == id) + .unwrap(); + assert_eq!(managed.stored, 4); + assert_eq!(managed.state, ManagedArtifactState::Partial); + + delete_managed_artifact(id, &models_path).unwrap(); + assert!(!partial.exists()); + fs::remove_dir_all(models_path).unwrap(); + } + + #[test] + fn restart_resumes_at_the_existing_partial_byte() { + let content = b"restart-resume works"; + let offset = 8; + let directory = std::env::temp_dir().join(format!( + "ds4-server-resume-{}", + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_nanos() + )); + fs::create_dir_all(&directory).unwrap(); + let partial = directory.join("model.gguf.part"); + fs::write(&partial, &content[..offset]).unwrap(); + + let listener = match TcpListener::bind("127.0.0.1:0") { + Ok(listener) => listener, + Err(error) if error.kind() == std::io::ErrorKind::PermissionDenied => { + fs::remove_dir_all(directory).unwrap(); + return; + } + Err(error) => panic!("could not start test server: {error}"), + }; + let address = listener.local_addr().unwrap(); + let server = thread::spawn(move || { + let (mut connection, _) = listener.accept().unwrap(); + let mut request = [0; 2048]; + let count = connection.read(&mut request).unwrap(); + let request = String::from_utf8_lossy(&request[..count]).to_ascii_lowercase(); + assert!(request.contains("range: bytes=8-")); + let remaining = &content[offset..]; + write!( + connection, + "HTTP/1.1 206 Partial Content\r\nContent-Length: {}\r\nContent-Range: bytes {offset}-{}/{}\r\nConnection: close\r\n\r\n", + remaining.len(), + content.len() - 1, + content.len(), + ) + .unwrap(); + connection.write_all(remaining).unwrap(); + }); + + let outcome = download_url_to_partial( + &format!("http://{address}/model.gguf"), + &partial, + &AtomicBool::new(false), + ) + .unwrap(); + server.join().unwrap(); + assert_eq!(outcome, DownloadOutcome::Complete); + assert_eq!(fs::read(&partial).unwrap(), content); + fs::remove_dir_all(directory).unwrap(); + } + + #[test] + fn cancellation_keeps_the_partial_file_for_the_next_run() { + let directory = std::env::temp_dir().join(format!( + "ds4-server-cancel-{}", + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap() + .as_nanos() + )); + let artifact = Artifact { + label: "test model", + file_name: "model.gguf", + repository: "unused", + size: 10, + sha256: "unused", + support: None, + }; + let partial = artifact.partial_path(ModelChoice::DeepSeekV4Flash, &directory); + fs::create_dir_all(partial.parent().unwrap()).unwrap(); + fs::write(&partial, b"part").unwrap(); + let cancel = AtomicBool::new(true); + + assert_eq!( + download_artifact_with_cancel( + ModelChoice::DeepSeekV4Flash, + &artifact, + &directory, + &cancel, + &AtomicU64::new(0), + ) + .unwrap(), + DownloadOutcome::Stopped + ); + assert_eq!(fs::read(&partial).unwrap(), b"part"); + fs::remove_dir_all(directory).unwrap(); + } +} diff --git a/src/server.rs b/src/server.rs index b75a820..aab41b8 100644 --- a/src/server.rs +++ b/src/server.rs @@ -1,3 +1,25 @@ +mod http; +mod request; +mod response; +mod tools; + +use http::{ + random_id, random_tool_id, read_request, send_error, send_json, send_response, send_sse, + send_sse_error, send_sse_headers, unix_time, +}; +use request::{ + anthropic_request, completion_request, parse_chat_request, raw_tool_schemas, responses_request, +}; +use response::{final_response, stream_response}; +#[cfg(test)] +use response::{receive_stream_event, send_stream_start}; +use tools::{ + TOOL_SYNTAXES, ToolProjectionEvent, ToolProjector, content_text, parse_generated_tools, + parse_generated_tools_with_ids, render_messages, tool_calls_json, +}; +#[cfg(test)] +use tools::{canonical_tools, validate_tool_results}; + use crate::database::AppPreferences; use crate::engine::ChatTurn; use crate::metrics::Metrics; @@ -163,304 +185,6 @@ enum Protocol { Responses, } -enum ToolProjectionState { - Seeking, - Invokes, - Parameters, - Value, - Done, - Failed, -} - -enum ToolProjectionEvent { - Text(String), - Start { - index: usize, - id: String, - name: String, - }, - Arguments { - index: usize, - fragment: String, - }, - End { - index: usize, - }, -} - -struct ToolProjector { - raw: String, - position: usize, - text_emitted: usize, - state: ToolProjectionState, - index: usize, - ids: Vec, - first_parameter: bool, - string_parameter: bool, - syntax: Option, -} - -impl ToolProjector { - fn new() -> Self { - Self { - raw: String::new(), - position: 0, - text_emitted: 0, - state: ToolProjectionState::Seeking, - index: 0, - ids: Vec::new(), - first_parameter: true, - string_parameter: false, - syntax: None, - } - } - - fn push(&mut self, chunk: &str, final_chunk: bool, prefix: &str) -> Vec { - self.raw.push_str(chunk); - let mut events = Vec::new(); - loop { - match self.state { - ToolProjectionState::Seeking => { - if let Some((start, syntax)) = TOOL_SYNTAXES - .iter() - .filter_map(|syntax| { - self.raw - .find(syntax.tool_start) - .map(|start| (start, *syntax)) - }) - .min_by_key(|(start, _)| *start) - { - if start > self.text_emitted { - let text = self.raw[self.text_emitted..start].trim_end(); - if !text.is_empty() { - events.push(ToolProjectionEvent::Text(text.to_owned())); - } - } - self.position = start + syntax.tool_start.len(); - self.text_emitted = start; - self.syntax = Some(syntax); - self.state = ToolProjectionState::Invokes; - } else { - let limit = if final_chunk { - self.raw.len() - } else { - TOOL_SYNTAXES - .iter() - .map(|syntax| { - safe_before_partial_marker(&self.raw, syntax.tool_start) - }) - .min() - .unwrap_or(self.raw.len()) - }; - if limit > self.text_emitted { - let text = &self.raw[self.text_emitted..limit]; - if !text.trim().is_empty() { - events.push(ToolProjectionEvent::Text(text.to_owned())); - self.text_emitted = limit; - } - } - break; - } - } - ToolProjectionState::Invokes => { - let syntax = self.syntax.unwrap(); - self.skip_whitespace(); - if self.full_at(syntax.tool_end) { - self.position += syntax.tool_end.len(); - self.state = ToolProjectionState::Done; - break; - } - if self.partial_at(syntax.tool_end) || self.partial_at(syntax.invoke_start) { - break; - } - if !self.full_at(syntax.invoke_start) { - self.state = ToolProjectionState::Failed; - break; - } - let Some(tag_end) = self.raw[self.position..].find('>') else { - break; - }; - let tag_end = self.position + tag_end + 1; - let Some(name) = dsml_attribute(&self.raw[self.position..tag_end], "name") - else { - self.state = ToolProjectionState::Failed; - break; - }; - let id = random_tool_id(prefix); - self.ids.push(id.clone()); - events.push(ToolProjectionEvent::Start { - index: self.index, - id, - name, - }); - events.push(ToolProjectionEvent::Arguments { - index: self.index, - fragment: "{".into(), - }); - self.position = tag_end; - self.first_parameter = true; - self.state = ToolProjectionState::Parameters; - } - ToolProjectionState::Parameters => { - let syntax = self.syntax.unwrap(); - self.skip_whitespace(); - if self.full_at(syntax.invoke_end) { - events.push(ToolProjectionEvent::Arguments { - index: self.index, - fragment: "}".into(), - }); - events.push(ToolProjectionEvent::End { index: self.index }); - self.position += syntax.invoke_end.len(); - self.index += 1; - self.state = ToolProjectionState::Invokes; - continue; - } - if self.partial_at(syntax.invoke_end) || self.partial_at(syntax.parameter_start) - { - break; - } - if !self.full_at(syntax.parameter_start) { - self.state = ToolProjectionState::Failed; - break; - } - let Some(tag_end) = self.raw[self.position..].find('>') else { - break; - }; - let tag_end = self.position + tag_end + 1; - let tag = &self.raw[self.position..tag_end]; - let Some(name) = dsml_attribute(tag, "name") else { - self.state = ToolProjectionState::Failed; - break; - }; - self.string_parameter = - dsml_attribute(tag, "string").as_deref() != Some("false"); - let mut fragment = if self.first_parameter { - String::new() - } else { - ",".into() - }; - self.first_parameter = false; - fragment - .push_str(&serde_json::to_string(&name).unwrap_or_else(|_| "\"\"".into())); - fragment.push(':'); - if self.string_parameter { - fragment.push('"'); - } - events.push(ToolProjectionEvent::Arguments { - index: self.index, - fragment, - }); - self.position = tag_end; - self.state = ToolProjectionState::Value; - } - ToolProjectionState::Value => { - let syntax = self.syntax.unwrap(); - if let Some(relative_end) = self.raw[self.position..].find(syntax.parameter_end) - { - let end = self.position + relative_end; - self.emit_value(end, &mut events); - if self.string_parameter { - events.push(ToolProjectionEvent::Arguments { - index: self.index, - fragment: "\"".into(), - }); - } - self.position = end + syntax.parameter_end.len(); - self.state = ToolProjectionState::Parameters; - continue; - } - let limit = safe_parameter_value_limit( - &self.raw, - self.position, - syntax.parameter_end, - self.string_parameter, - ); - self.emit_value(limit, &mut events); - break; - } - ToolProjectionState::Done | ToolProjectionState::Failed => break, - } - } - events - } - - fn emit_value(&mut self, end: usize, events: &mut Vec) { - if end <= self.position { - return; - } - let raw = &self.raw[self.position..end]; - let fragment = if self.string_parameter { - let value = unescape_dsml(raw); - let encoded = serde_json::to_string(&value).unwrap_or_else(|_| "\"\"".into()); - encoded[1..encoded.len() - 1].to_owned() - } else { - raw.to_owned() - }; - events.push(ToolProjectionEvent::Arguments { - index: self.index, - fragment, - }); - self.position = end; - } - - fn skip_whitespace(&mut self) { - while self.raw[self.position..] - .chars() - .next() - .is_some_and(char::is_whitespace) - { - self.position += self.raw[self.position..].chars().next().unwrap().len_utf8(); - } - } - - fn full_at(&self, marker: &str) -> bool { - self.raw.as_bytes()[self.position..].starts_with(marker.as_bytes()) - } - - fn partial_at(&self, marker: &str) -> bool { - let tail = &self.raw.as_bytes()[self.position..]; - tail.len() < marker.len() && marker.as_bytes().starts_with(tail) - } -} - -fn dsml_attribute(tag: &str, name: &str) -> Option { - let start = tag.find(&format!("{name}=\""))? + name.len() + 2; - let end = start + tag[start..].find('"')?; - Some(unescape_dsml(&tag[start..end])) -} - -fn safe_before_partial_marker(text: &str, marker: &str) -> usize { - let mut limit = text.len().saturating_sub(marker.len().saturating_sub(1)); - while !text.is_char_boundary(limit) { - limit -= 1; - } - limit -} - -fn safe_parameter_value_limit(text: &str, start: usize, end_marker: &str, string: bool) -> usize { - let bytes = text.as_bytes(); - let marker = end_marker.as_bytes(); - let mut limit = bytes.len(); - for length in (1..marker.len().min(bytes.len().saturating_sub(start) + 1)).rev() { - if bytes[start..].ends_with(&marker[..length]) { - limit -= length; - break; - } - } - if string { - for entity in ["&", "<", ">", """, "'"] { - let entity = entity.as_bytes(); - for length in 1..entity.len() { - if bytes[start..limit].ends_with(&entity[..length]) { - limit -= length; - return limit; - } - } - } - } - limit -} - struct HttpRequest { method: String, path: String, @@ -727,1935 +451,6 @@ fn compatible_completion( } } -fn completion_request(mut value: Value) -> Result { - let object = value - .as_object_mut() - .ok_or_else(|| (400, "invalid JSON request".to_owned()))?; - let prompt = object - .remove("prompt") - .ok_or_else(|| (400, "missing prompt".to_owned()))?; - let prompt = match prompt { - Value::String(text) => text, - Value::Array(values) => values - .into_iter() - .next() - .and_then(|value| value.as_str().map(str::to_owned)) - .unwrap_or_default(), - _ => String::new(), - }; - object.insert( - "messages".into(), - json!([ - {"role": "system", "content": "You are a helpful assistant"}, - {"role": "user", "content": prompt} - ]), - ); - decode_chat_request(value) -} - -fn anthropic_request(value: Value) -> Result { - let object = value - .as_object() - .ok_or_else(|| (400, "invalid JSON request".to_owned()))?; - let source = object - .get("messages") - .and_then(Value::as_array) - .ok_or_else(|| (400, "missing messages".to_owned()))?; - let mut messages = Vec::new(); - let system = object.get("system").map(content_text).unwrap_or_default(); - if !system.is_empty() { - messages.push(json!({"role": "system", "content": system})); - } - for message in source { - let role = message - .get("role") - .and_then(Value::as_str) - .unwrap_or("user"); - let Some(blocks) = message.get("content").and_then(Value::as_array) else { - messages.push(json!({"role": role, "content": message.get("content").cloned().unwrap_or(Value::Null)})); - continue; - }; - let mut text = String::new(); - let mut reasoning = String::new(); - let mut calls = Vec::new(); - for block in blocks { - match block.get("type").and_then(Value::as_str).unwrap_or("text") { - "text" => text.push_str(block.get("text").and_then(Value::as_str).unwrap_or("")), - "thinking" | "redacted_thinking" => reasoning.push_str( - block - .get("thinking") - .or_else(|| block.get("data")) - .and_then(Value::as_str) - .unwrap_or(""), - ), - "tool_use" => calls.push(json!({ - "id": block.get("id").and_then(Value::as_str).unwrap_or(""), - "type": "function", - "function": { - "name": block.get("name").and_then(Value::as_str).unwrap_or(""), - "arguments": block.get("input").cloned().unwrap_or_else(|| json!({})).to_string() - } - })), - "tool_result" => { - if !text.is_empty() { - messages.push(json!({"role": role, "content": std::mem::take(&mut text)})); - } - messages.push(json!({ - "role": "tool", - "content": block.get("content").map(content_text).unwrap_or_default(), - "tool_call_id": block.get("tool_use_id").and_then(Value::as_str).unwrap_or("") - })); - } - _ => {} - } - } - if !text.is_empty() || !reasoning.is_empty() || !calls.is_empty() { - messages.push(json!({ - "role": role, - "content": text, - "reasoning_content": reasoning, - "tool_calls": calls - })); - } - } - let mut request = Map::new(); - copy_request_fields(object, &mut request); - request.insert("messages".into(), Value::Array(messages)); - if let Some(stops) = object.get("stop_sequences") { - request.insert("stop".into(), stops.clone()); - } - if let Some(effort) = object - .get("output_config") - .and_then(|value| value.get("effort")) - { - request.insert("reasoning_effort".into(), effort.clone()); - } - request.insert( - "tools".into(), - normalize_anthropic_tools(object.get("tools")), - ); - decode_chat_request(Value::Object(request)) -} - -fn responses_request(value: Value) -> Result { - let object = value - .as_object() - .ok_or_else(|| (400, "invalid JSON request".to_owned()))?; - for key in ["previous_response_id", "conversation"] { - if object.get(key).is_some_and(|value| !value.is_null()) { - return Err(( - 400, - format!("{key} is not supported; replay full input instead"), - )); - } - } - if let Some(choice) = object.get("tool_choice") { - match choice { - Value::String(choice) if choice == "none" || choice == "auto" => {} - Value::String(choice) => { - return Err((400, format!("tool_choice={choice} not supported"))); - } - Value::Object(_) => return Err((400, "forced tool_choice not supported".into())), - _ => {} - } - } - let input = object - .get("input") - .ok_or_else(|| (400, "missing input".to_owned()))?; - let mut messages = Vec::new(); - match input { - Value::String(text) => messages.push(json!({"role": "user", "content": text})), - Value::Array(items) => { - let mut pending_reasoning = String::new(); - for item in items { - let item = item - .as_object() - .ok_or_else(|| (400, "invalid JSON request".to_owned()))?; - if item - .get("status") - .and_then(Value::as_str) - .is_some_and(|status| status != "completed") - { - return Err((400, "invalid JSON request".into())); - } - let kind = item - .get("type") - .and_then(Value::as_str) - .unwrap_or("message"); - let role = item.get("role").and_then(Value::as_str).unwrap_or("user"); - let consumes_reasoning = (kind == "message" && role == "assistant") - || matches!( - kind, - "function_call" - | "custom_tool_call" - | "local_shell_call" - | "web_search_call" - | "tool_search_call" - | "image_generation_call" - ); - let bookkeeping = matches!(kind, "compaction" | "context_compaction"); - if !consumes_reasoning && !bookkeeping && !pending_reasoning.is_empty() { - messages.push(reasoning_message(std::mem::take(&mut pending_reasoning))); - } - match kind { - "message" => { - let mut message = json!({ - "role": role, - "content": responses_content_text(item.get("content").unwrap_or(&Value::Null))? - }); - if role == "assistant" && !pending_reasoning.is_empty() { - message["reasoning_content"] = - Value::String(std::mem::take(&mut pending_reasoning)); - } - messages.push(message); - } - "function_call" | "custom_tool_call" => { - let name = if kind == "function_call" { - format!( - "{}{}", - item.get("namespace").and_then(Value::as_str).unwrap_or(""), - item.get("name").and_then(Value::as_str).unwrap_or("") - ) - } else { - item.get("name") - .and_then(Value::as_str) - .unwrap_or("") - .to_owned() - }; - push_responses_call( - &mut messages, - &mut pending_reasoning, - item, - &name, - item.get("arguments").or_else(|| item.get("input")), - ); - } - "function_call_output" | "custom_tool_call_output" => messages.push(json!({ - "role": "tool", - "content": responses_output_text(item.get("output"))?, - "tool_call_id": item.get("call_id").or_else(|| item.get("id")).and_then(Value::as_str).unwrap_or("") - })), - "reasoning" => { - for value in [item.get("summary"), item.get("content")] - .into_iter() - .flatten() - { - let text = responses_content_text(value)?; - if !pending_reasoning.is_empty() && !text.is_empty() { - pending_reasoning.push('\n'); - } - pending_reasoning.push_str(&text); - } - } - "local_shell_call" - | "web_search_call" - | "tool_search_call" - | "image_generation_call" => { - let name = match kind { - "local_shell_call" => "local_shell", - "tool_search_call" => "tool_search", - other => other, - }; - push_responses_call( - &mut messages, - &mut pending_reasoning, - item, - name, - item.get("action") - .or_else(|| item.get("arguments")) - .or_else(|| item.get("input")), - ); - } - "local_shell_call_output" - | "web_search_call_output" - | "tool_search_output" - | "tool_search_call_output" - | "image_generation_call_output" => { - let content = if let Some(value) = - item.get("output").or_else(|| item.get("result")) - { - responses_output_text(Some(value))? - } else { - item.get("tools").map(Value::to_string).unwrap_or_default() - }; - messages.push(json!({ - "role": "tool", - "content": content, - "tool_call_id": item.get("call_id").or_else(|| item.get("id")).and_then(Value::as_str).unwrap_or("") - })); - } - "compaction" | "context_compaction" => {} - _ => return Err((400, "invalid JSON request".into())), - } - } - if !pending_reasoning.is_empty() { - messages.push(reasoning_message(pending_reasoning)); - } - } - _ => return Err((400, "invalid JSON request".into())), - } - if let Some(instructions) = object.get("instructions") { - let instructions = match instructions { - Value::Null => String::new(), - Value::String(text) => text.clone(), - _ => return Err((400, "invalid JSON request".into())), - }; - if !instructions.is_empty() { - messages.insert(0, json!({"role": "system", "content": instructions})); - } - } - let mut request = Map::new(); - copy_request_fields(object, &mut request); - request.insert("messages".into(), Value::Array(messages)); - if let Some(tokens) = object.get("max_output_tokens") { - request.insert("max_tokens".into(), tokens.clone()); - } - if let Some(effort) = object - .get("reasoning") - .and_then(|value| value.get("effort")) - { - request.insert("reasoning_effort".into(), effort.clone()); - } - request.insert( - "tools".into(), - normalize_responses_tools(object.get("tools")), - ); - decode_chat_request(Value::Object(request)) -} - -fn reasoning_message(reasoning: String) -> Value { - json!({"role": "assistant", "content": "", "reasoning_content": reasoning}) -} - -fn push_responses_call( - messages: &mut Vec, - pending_reasoning: &mut String, - item: &Map, - name: &str, - arguments: Option<&Value>, -) { - let call = json!({ - "id": item.get("call_id").or_else(|| item.get("id")).and_then(Value::as_str).unwrap_or(""), - "type": "function", - "function": {"name": name, "arguments": raw_json_text(arguments)} - }); - if let Some(last) = messages.last_mut() - && last.get("role").and_then(Value::as_str) == Some("assistant") - { - if !pending_reasoning.is_empty() - && last - .get("reasoning_content") - .and_then(Value::as_str) - .is_none_or(str::is_empty) - { - last["reasoning_content"] = Value::String(std::mem::take(pending_reasoning)); - } - if !last.get("tool_calls").is_some_and(Value::is_array) { - last["tool_calls"] = json!([]); - } - last["tool_calls"].as_array_mut().unwrap().push(call); - } else { - messages.push(json!({ - "role": "assistant", - "content": "", - "reasoning_content": std::mem::take(pending_reasoning), - "tool_calls": [call] - })); - } -} - -fn raw_json_text(value: Option<&Value>) -> String { - match value { - Some(Value::String(text)) => text.clone(), - Some(value) => value.to_string(), - None => "{}".into(), - } -} - -fn responses_output_text(value: Option<&Value>) -> Result { - match value { - Some(Value::Array(_)) => responses_content_text(value.unwrap()), - Some(Value::String(text)) => Ok(text.clone()), - Some(value) => Ok(value.to_string()), - None => Ok(String::new()), - } -} - -fn responses_content_text(value: &Value) -> Result { - match value { - Value::String(text) => Ok(text.clone()), - Value::Null => Ok(String::new()), - Value::Array(parts) => { - let mut text = String::new(); - for part in parts { - match part { - Value::String(part) => text.push_str(part), - Value::Object(part) - if matches!( - part.get("type").and_then(Value::as_str), - Some( - "input_text" - | "output_text" - | "text" - | "summary_text" - | "reasoning_text" - ) - ) && matches!( - part.get("text"), - Some(Value::String(_) | Value::Null) - ) => - { - if let Some(value) = part.get("text").and_then(Value::as_str) { - text.push_str(value); - } - } - _ => return Err((400, "invalid JSON request".into())), - } - } - Ok(text) - } - _ => Err((400, "invalid JSON request".into())), - } -} - -fn copy_request_fields(source: &Map, target: &mut Map) { - for key in [ - "model", - "max_tokens", - "max_completion_tokens", - "temperature", - "top_p", - "min_p", - "top_k", - "seed", - "stream", - "stream_options", - "thinking", - "think", - "reasoning_effort", - "tool_choice", - "stop", - ] { - if let Some(value) = source.get(key) { - target.insert(key.into(), value.clone()); - } - } -} - -fn normalize_anthropic_tools(tools: Option<&Value>) -> Value { - Value::Array( - tools - .and_then(Value::as_array) - .into_iter() - .flatten() - .map(|tool| { - json!({"type": "function", "function": { - "name": tool.get("name").and_then(Value::as_str).unwrap_or(""), - "description": tool.get("description").and_then(Value::as_str).unwrap_or(""), - "parameters": tool.get("input_schema").cloned().unwrap_or_else(|| json!({"type": "object"})) - }}) - }) - .collect(), - ) -} - -fn normalize_responses_tools(tools: Option<&Value>) -> Value { - Value::Array( - tools - .and_then(Value::as_array) - .into_iter() - .flatten() - .filter(|tool| tool.get("type").and_then(Value::as_str) == Some("function")) - .map(|tool| { - json!({"type": "function", "function": { - "name": tool.get("name").and_then(Value::as_str).unwrap_or(""), - "description": tool.get("description").and_then(Value::as_str).unwrap_or(""), - "parameters": tool.get("parameters").cloned().unwrap_or_else(|| json!({"type": "object"})) - }}) - }) - .collect(), - ) -} - -fn decode_chat_request(value: Value) -> Result { - serde_json::from_value(value).map_err(|_| (400, "invalid JSON request".to_owned())) -} - -fn raw_tool_schemas(body: &[u8], protocol: Protocol) -> Result, (u16, String)> { - let request: RawToolsRequest<'_> = - serde_json::from_slice(body).map_err(|_| (400, "invalid JSON request".to_owned()))?; - request - .tools - .into_iter() - .map(|tool| { - if protocol == Protocol::Chat { - let wrapper: RawOpenAiTool<'_> = serde_json::from_str(tool.get()) - .map_err(|_| (400, "invalid JSON request".to_owned()))?; - Ok(wrapper.function.unwrap_or(tool).get().to_owned()) - } else { - Ok(tool.get().to_owned()) - } - }) - .collect() -} - -fn parse_chat_request( - state: &State, - request: ChatRequest, - protocol: Protocol, -) -> Result { - if request.messages.is_empty() { - return Err((400, "missing messages".into())); - } - let requested_id = request.model.as_deref().unwrap_or_default(); - let preferences = state - .preferences - .read() - .map_err(|_| (500, "preferences are unavailable".to_owned()))? - .clone(); - let model = if requested_id.is_empty() { - ModelChoice::from_id(&preferences.selected_model) - } else { - model_alias(requested_id) - } - .ok_or_else(|| (400, format!("unknown model: {requested_id}")))?; - if !installed_endpoint_models(&state.models_path).contains(&model) { - return Err((400, format!("model is not installed and verified: {model}"))); - } - - let mut generation = preferences.generation().map_err(|error| (500, error))?; - generation.system_prompt.clear(); - generation.max_generated_tokens = request - .max_completion_tokens - .or(request.max_tokens) - .unwrap_or(generation.max_generated_tokens); - if generation.max_generated_tokens <= 0 { - return Err((400, "max_tokens must be positive".into())); - } - generation.temperature = request.temperature.or(Some(1.0)); - generation.top_p = request.top_p.or(Some(1.0)); - generation.min_p = request.min_p.or(Some(0.05)); - generation.seed = request.seed.filter(|seed| *seed > 0); - generation.reasoning_mode = request_reasoning(&request, requested_id)?; - - let tools_enabled = (!request.tools.is_empty() || !request.tool_schemas.is_empty()) - && request.tool_choice.as_ref().and_then(Value::as_str) != Some("none"); - let (system, messages) = render_messages( - state, - &request.messages, - &request.tools, - &request.tool_schemas, - tools_enabled, - protocol, - )?; - generation.system_prompt = system; - let runtime = preferences.runtime().map_err(|error| (500, error))?; - let mut effective = effective_settings(model, &generation, &runtime, &state.models_path) - .map_err(|error| (400, error))?; - effective.turn.top_k = request.top_k.unwrap_or(0); - if effective.turn.top_k < 0 { - return Err((400, "top_k must not be negative".into())); - } - effective.turn.stops = match request.stop { - Some(OneOrMany::One(stop)) => vec![stop], - Some(OneOrMany::Many(stops)) => stops, - None => Vec::new(), - }; - effective.turn.stops.retain(|stop| !stop.is_empty()); - - Ok(ParsedRequest { - protocol, - model_id: if requested_id.is_empty() { - model.id().to_owned() - } else { - requested_id.to_owned() - }, - messages, - turn: effective.turn, - engine: effective.engine, - idle_timeout: Duration::from_secs(preferences.idle_timeout_minutes.max(1) as u64 * 60), - stream: request.stream, - include_usage: request - .stream_options - .is_some_and(|options| options.include_usage), - has_tools: tools_enabled, - }) -} - -fn request_reasoning( - request: &ChatRequest, - model_id: &str, -) -> Result { - let explicit_thinking = request - .think - .or_else(|| request.thinking.as_ref().and_then(thinking_enabled)); - let mut reasoning = match request.reasoning_effort.as_deref() { - Some("max") => ReasoningMode::Max, - Some("none") => ReasoningMode::Direct, - Some("xhigh" | "high" | "medium" | "low" | "minimal") | None => ReasoningMode::High, - Some(value) => return Err((400, format!("unsupported reasoning_effort: {value}"))), - }; - if explicit_thinking == Some(false) - || (explicit_thinking.is_none() - && matches!( - model_id, - "deepseek-chat" | "glm-5.2-chat" | "glm-5.2-no-think" | "glm-5.2-nothink" - )) - { - reasoning = ReasoningMode::Direct; - } - Ok(reasoning) -} - -fn thinking_enabled(value: &Value) -> Option { - value.as_bool().or_else(|| { - value.as_str().map(|value| value != "disabled").or_else(|| { - value - .get("type") - .and_then(Value::as_str) - .map(|value| value != "disabled") - }) - }) -} - -fn render_messages( - state: &State, - messages: &[ApiMessage], - tools: &[Value], - tool_schemas: &[String], - tools_enabled: bool, - protocol: Protocol, -) -> Result<(String, Vec), (u16, String)> { - validate_tool_results(state, messages, protocol)?; - let preserve_reasoning = tools_enabled - || messages.iter().any(|message| { - matches!(message.role.as_str(), "tool" | "function") || !message.tool_calls.is_empty() - }); - let mut system = String::new(); - if tools_enabled { - system.push_str(TOOLS_PROMPT); - if tool_schemas.is_empty() { - for tool in tools { - let schema = tool.get("function").unwrap_or(tool); - if !system.ends_with("\n\n") { - system.push('\n'); - } - system.push_str( - &serde_json::to_string(schema) - .map_err(|error| (400, format!("invalid tool schema: {error}")))?, - ); - system.push('\n'); - } - } else { - for schema in tool_schemas { - if !system.ends_with("\n\n") { - system.push('\n'); - } - system.push_str(schema); - system.push('\n'); - } - } - system.push_str( - "\nYou MUST strictly follow the above defined tool name and parameter schemas to invoke tool calls. Use the exact parameter names from the schemas.", - ); - } - - let mut turns = Vec::::new(); - for message in messages { - let content = content_text(&message.content); - match message.role.as_str() { - "system" | "developer" => { - if !system.is_empty() { - system.push_str("\n\n"); - } - system.push_str(&content); - } - "user" => turns.push(ChatTurn { - user: true, - skip_previous_eos: false, - reasoning: None, - reasoning_complete: true, - content, - }), - "tool" | "function" => { - let wrapped = format!( - "{}", - escape_tool_result(&content) - ); - if let Some(previous) = turns.last_mut() - && previous.user - && previous.content.starts_with("") - { - previous.content.push_str(&wrapped); - } else { - turns.push(ChatTurn { - user: true, - skip_previous_eos: protocol == Protocol::Responses, - reasoning: None, - reasoning_complete: true, - content: wrapped, - }); - } - } - "assistant" => { - let mut content = content; - if !message.tool_calls.is_empty() { - content.push_str(&replayed_or_canonical_tools(state, &message.tool_calls)); - } - let reasoning = content_text(&message.reasoning_content); - turns.push(ChatTurn { - user: false, - skip_previous_eos: false, - reasoning: (preserve_reasoning && !reasoning.is_empty()).then_some(reasoning), - reasoning_complete: true, - content, - }); - } - role => return Err((400, format!("unsupported message role: {role}"))), - } - } - Ok((system, turns)) -} - -fn validate_tool_results( - state: &State, - messages: &[ApiMessage], - protocol: Protocol, -) -> Result<(), (u16, String)> { - if !matches!(protocol, Protocol::Anthropic | Protocol::Responses) { - return Ok(()); - } - let memory = state.tool_memory.lock().ok(); - for (index, message) in messages.iter().enumerate() { - if !matches!(message.role.as_str(), "tool" | "function") || message.tool_call_id.is_empty() - { - continue; - } - let id = &message.tool_call_id; - let live = memory - .as_ref() - .is_some_and(|memory| memory.contains_key(id)); - let replayed = messages[..index].iter().any(|message| { - message.role == "assistant" && message.tool_calls.iter().any(|call| call.id == *id) - }); - if live || replayed { - continue; - } - let message = match protocol { - Protocol::Anthropic => format!( - "Anthropic continuation state is not available for tool_use_id {id}; retry by replaying the full messages history" - ), - Protocol::Responses => format!( - "Responses continuation state is not available for call_id {id}; retry by replaying the full input history" - ), - _ => unreachable!(), - }; - return Err((400, message)); - } - Ok(()) -} - -fn replayed_or_canonical_tools(state: &State, calls: &[ApiToolCall]) -> String { - if let Ok(memory) = state.tool_memory.lock() - && let Some(raw) = calls.iter().find_map(|call| { - (!call.id.is_empty()) - .then(|| memory.get(&call.id)) - .flatten() - }) - { - return raw.clone(); - } - canonical_tools(calls) -} - -fn canonical_tools(calls: &[ApiToolCall]) -> String { - let mut output = String::from("\n\n<|DSML|tool_calls>\n"); - for call in calls { - output.push_str("<|DSML|invoke name=\""); - output.push_str(&escape_attribute(&call.function.name)); - output.push_str("\">\n"); - match serde_json::from_str::(&call.function.arguments) { - Ok(Value::Object(arguments)) => { - for (name, value) in arguments { - output.push_str("<|DSML|parameter name=\""); - output.push_str(&escape_attribute(&name)); - let string = value.as_str(); - output.push_str(if string.is_some() { - "\" string=\"true\">" - } else { - "\" string=\"false\">" - }); - if let Some(value) = string { - output.push_str(&escape_parameter(value)); - } else { - output.push_str(&escape_json_parameter(&value.to_string())); - } - output.push_str("\n"); - } - } - _ => { - output.push_str("<|DSML|parameter name=\"arguments\" string=\"true\">"); - output.push_str(&escape_parameter(&call.function.arguments)); - output.push_str("\n"); - } - } - output.push_str("\n"); - } - output.push_str(""); - output -} - -fn final_response( - stream: &mut TcpStream, - state: &State, - request: ResponseOptions, - active: crate::runtime::ActiveGeneration, - id: &str, -) -> Result<(), (u16, String)> { - let output = wait_for_output(active)?; - let (content, calls) = parse_generated_tools(state, &output.message.content, request.protocol); - let finish = if calls.is_empty() { - output.finish_reason - } else { - "tool_calls" - }; - let reasoning = output.message.reasoning.filter(|value| !value.is_empty()); - let usage = usage_json( - output.prompt_tokens, - output.cached_tokens, - output.completion_tokens, - ); - let body = match request.protocol { - Protocol::Chat => { - let mut message = json!({"role": "assistant", "content": content}); - if let Some(reasoning) = reasoning { - message["reasoning_content"] = Value::String(reasoning); - } - if !calls.is_empty() { - message["tool_calls"] = tool_calls_json(&calls); - } - json!({ - "id": id, "object": "chat.completion", "created": unix_time(), - "model": request.model_id, - "choices": [{"index": 0, "message": message, "finish_reason": finish}], - "usage": usage, - }) - } - Protocol::Completion => json!({ - "id": id, "object": "text_completion", "created": unix_time(), - "model": request.model_id, - "choices": [{"text": content, "index": 0, "finish_reason": finish}], - "usage": usage, - }), - Protocol::Anthropic => { - let mut blocks = Vec::new(); - if let Some(reasoning) = reasoning { - blocks.push(json!({"type": "thinking", "thinking": reasoning, "signature": id})); - } - if !content.is_empty() { - blocks.push(json!({"type": "text", "text": content})); - } - for call in &calls { - blocks.push(json!({ - "type": "tool_use", "id": call.id, "name": call.function.name, - "input": serde_json::from_str::(&call.function.arguments).unwrap_or_else(|_| json!({})) - })); - } - if blocks.is_empty() || (blocks.iter().all(|block| block["type"] == "thinking")) { - blocks.push(json!({"type": "text", "text": ""})); - } - let cached = output.cached_tokens.min(output.prompt_tokens); - let written = output.prompt_tokens - cached; - json!({ - "id": id, "type": "message", "role": "assistant", "model": request.model_id, - "content": blocks, - "stop_reason": if finish == "tool_calls" { "tool_use" } else if finish == "length" { "max_tokens" } else { "end_turn" }, - "stop_sequence": Value::Null, - "usage": { - "input_tokens": output.prompt_tokens - cached - written, - "output_tokens": output.completion_tokens, - "cache_read_input_tokens": cached, - "cache_creation_input_tokens": written - } - }) - } - Protocol::Responses => { - let status = if finish == "length" { - "incomplete" - } else if finish == "error" { - "failed" - } else { - "completed" - }; - let mut items = Vec::new(); - if let Some(reasoning) = reasoning { - items.push(json!({ - "id": random_id("rs_"), "type": "reasoning", "status": status, - "summary": [{"type": "summary_text", "text": reasoning}] - })); - } - if !content.is_empty() { - items.push(json!({ - "id": random_id("msg_"), "type": "message", "status": status, - "role": "assistant", "content": [{"type": "output_text", "text": content, "annotations": []}] - })); - } - for call in &calls { - items.push(json!({ - "id": random_id("fc_"), "type": "function_call", "status": status, - "name": call.function.name, "call_id": call.id, "arguments": call.function.arguments - })); - } - json!({ - "id": id, "object": "response", "created_at": unix_time(), "status": status, - "model": request.model_id, "output": items, - "usage": { - "input_tokens": output.prompt_tokens, - "input_tokens_details": {"cached_tokens": output.cached_tokens.min(output.prompt_tokens), "cache_write_tokens": output.prompt_tokens - output.cached_tokens.min(output.prompt_tokens)}, - "output_tokens": output.completion_tokens, - "output_tokens_details": {"reasoning_tokens": 0}, - "total_tokens": output.prompt_tokens + output.completion_tokens - } - }) - } - }; - send_json(stream, 200, &body).map_err(|error| (500, error)) -} - -fn stream_response( - stream: &mut TcpStream, - state: &State, - request: ResponseOptions, - active: crate::runtime::ActiveGeneration, - id: &str, -) -> Result<(), (u16, String)> { - if request.protocol != Protocol::Chat { - return structured_stream_response(stream, state, request, active, id); - } - stream_response_with_keepalive( - stream, - state, - request, - active, - id, - PREFILL_KEEPALIVE_INTERVAL, - ) -} - -fn structured_stream_response( - stream: &mut TcpStream, - state: &State, - request: ResponseOptions, - active: crate::runtime::ActiveGeneration, - id: &str, -) -> Result<(), (u16, String)> { - match request.protocol { - Protocol::Completion => completion_stream_response(stream, request, active, id), - Protocol::Anthropic => anthropic_stream_response(stream, state, request, active, id), - Protocol::Responses => responses_stream_response(stream, state, request, active, id), - Protocol::Chat => unreachable!(), - } -} - -fn completion_stream_response( - stream: &mut TcpStream, - request: ResponseOptions, - active: crate::runtime::ActiveGeneration, - id: &str, -) -> Result<(), (u16, String)> { - send_sse_headers(stream).map_err(|error| (500, error))?; - while let Ok(event) = active.events.recv() { - match event { - GenerationEvent::Chunk { - reasoning: false, - content, - } if !content.is_empty() => send_sse( - stream, - &json!({ - "id": id, "object": "text_completion", "created": unix_time(), - "model": request.model_id, - "choices": [{"text": content, "index": 0, "finish_reason": Value::Null}] - }), - ) - .map_err(|error| (500, error))?, - GenerationEvent::Finished(Ok(output)) => { - send_sse( - stream, - &json!({ - "id": id, "object": "text_completion", "created": unix_time(), - "model": request.model_id, - "choices": [{"text": "", "index": 0, "finish_reason": output.finish_reason}] - }), - ) - .map_err(|error| (500, error))?; - if request.include_usage { - send_sse( - stream, - &json!({ - "id": id, "object": "text_completion", "created": unix_time(), - "model": request.model_id, "choices": [], - "usage": usage_json(output.prompt_tokens, output.cached_tokens, output.completion_tokens) - }), - ) - .map_err(|error| (500, error))?; - } - return stream - .write_all(b"data: [DONE]\n\n") - .map_err(|error| (500, error.to_string())); - } - GenerationEvent::Finished(Err(error)) => { - let _ = send_sse_error(stream, &error); - return Ok(()); - } - _ => {} - } - } - Err((500, "The model runtime stopped unexpectedly.".into())) -} - -fn anthropic_stream_response( - stream: &mut TcpStream, - state: &State, - request: ResponseOptions, - active: crate::runtime::ActiveGeneration, - id: &str, -) -> Result<(), (u16, String)> { - send_sse_headers(stream).map_err(|error| (500, error))?; - let mut prompt_tokens = 0; - let mut started = false; - let mut block = None::<(usize, bool)>; - let mut next_index = 0; - let mut projector = ToolProjector::new(); - let mut tool_indices = Vec::new(); - while let Ok(event) = active.events.recv() { - match event { - GenerationEvent::Context { - used, - tokens_per_second, - .. - } => { - if tokens_per_second.is_none() { - prompt_tokens = used; - } else if !started { - anthropic_stream_start(stream, &request, id, prompt_tokens, 0)?; - started = true; - } - } - GenerationEvent::Chunk { reasoning, content } => { - if !started { - anthropic_stream_start(stream, &request, id, prompt_tokens, 0)?; - started = true; - } - if !reasoning && request.has_tools { - let events = projector.push(&content, false, "toolu_"); - send_anthropic_projection_events( - stream, - events, - &mut block, - &mut next_index, - &mut tool_indices, - )?; - continue; - } - if block.is_some_and(|(_, current_reasoning)| current_reasoning != reasoning) { - let (index, _) = block.take().unwrap(); - send_named_sse( - stream, - "content_block_stop", - &json!({"type": "content_block_stop", "index": index}), - )?; - } - let index = if let Some((index, _)) = block { - index - } else { - let index = next_index; - next_index += 1; - send_named_sse( - stream, - "content_block_start", - &json!({ - "type": "content_block_start", "index": index, - "content_block": if reasoning { json!({"type": "thinking", "thinking": "", "signature": ""}) } else { json!({"type": "text", "text": ""}) } - }), - )?; - block = Some((index, reasoning)); - index - }; - let delta = if reasoning { - json!({"type": "thinking_delta", "thinking": content}) - } else { - json!({"type": "text_delta", "text": content}) - }; - send_named_sse( - stream, - "content_block_delta", - &json!({"type": "content_block_delta", "index": index, "delta": delta}), - )?; - } - GenerationEvent::Finished(Ok(output)) => { - if !started { - anthropic_stream_start( - stream, - &request, - id, - output.prompt_tokens, - output.cached_tokens, - )?; - } - if request.has_tools { - let events = projector.push("", true, "toolu_"); - send_anthropic_projection_events( - stream, - events, - &mut block, - &mut next_index, - &mut tool_indices, - )?; - } - if let Some((index, _)) = block.take() { - send_named_sse( - stream, - "content_block_stop", - &json!({"type": "content_block_stop", "index": index}), - )?; - } - let (_, calls) = parse_generated_tools_with_ids( - state, - &output.message.content, - Protocol::Anthropic, - &projector.ids, - ); - if projector.ids.is_empty() { - for call in &calls { - send_named_sse( - stream, - "content_block_start", - &json!({"type": "content_block_start", "index": next_index, "content_block": {"type": "tool_use", "id": call.id, "name": call.function.name, "input": {}}}), - )?; - send_named_sse( - stream, - "content_block_delta", - &json!({"type": "content_block_delta", "index": next_index, "delta": {"type": "input_json_delta", "partial_json": call.function.arguments}}), - )?; - send_named_sse( - stream, - "content_block_stop", - &json!({"type": "content_block_stop", "index": next_index}), - )?; - next_index += 1; - } - } - let finish = if calls.is_empty() { - output.finish_reason - } else { - "tool_calls" - }; - send_named_sse( - stream, - "message_delta", - &json!({"type": "message_delta", "delta": {"stop_reason": if finish == "tool_calls" { "tool_use" } else if finish == "length" { "max_tokens" } else { "end_turn" }, "stop_sequence": Value::Null}, "usage": {"output_tokens": output.completion_tokens}}), - )?; - return send_named_sse(stream, "message_stop", &json!({"type": "message_stop"})); - } - GenerationEvent::Finished(Err(error)) => { - return send_named_sse( - stream, - "error", - &json!({"type": "error", "error": {"type": "api_error", "message": error}}), - ); - } - _ => {} - } - } - Err((500, "The model runtime stopped unexpectedly.".into())) -} - -fn send_anthropic_projection_events( - stream: &mut impl Write, - events: Vec, - block: &mut Option<(usize, bool)>, - next_index: &mut usize, - tool_indices: &mut Vec, -) -> Result<(), (u16, String)> { - for event in events { - match event { - ToolProjectionEvent::Text(text) if !text.is_empty() => { - if block.is_some_and(|(_, reasoning)| reasoning) { - let (index, _) = block.take().unwrap(); - send_named_sse( - stream, - "content_block_stop", - &json!({"type": "content_block_stop", "index": index}), - )?; - } - let index = if let Some((index, _)) = *block { - index - } else { - let index = *next_index; - *next_index += 1; - send_named_sse( - stream, - "content_block_start", - &json!({"type": "content_block_start", "index": index, "content_block": {"type": "text", "text": ""}}), - )?; - *block = Some((index, false)); - index - }; - send_named_sse( - stream, - "content_block_delta", - &json!({"type": "content_block_delta", "index": index, "delta": {"type": "text_delta", "text": text}}), - )?; - } - ToolProjectionEvent::Start { index, id, name } => { - if let Some((open_index, _)) = block.take() { - send_named_sse( - stream, - "content_block_stop", - &json!({"type": "content_block_stop", "index": open_index}), - )?; - } - let content_index = *next_index; - *next_index += 1; - if tool_indices.len() == index { - tool_indices.push(content_index); - } - send_named_sse( - stream, - "content_block_start", - &json!({"type": "content_block_start", "index": content_index, "content_block": {"type": "tool_use", "id": id, "name": name, "input": {}}}), - )?; - } - ToolProjectionEvent::Arguments { index, fragment } => { - if let Some(content_index) = tool_indices.get(index) { - send_named_sse( - stream, - "content_block_delta", - &json!({"type": "content_block_delta", "index": content_index, "delta": {"type": "input_json_delta", "partial_json": fragment}}), - )?; - } - } - ToolProjectionEvent::End { index } => { - if let Some(content_index) = tool_indices.get(index) { - send_named_sse( - stream, - "content_block_stop", - &json!({"type": "content_block_stop", "index": content_index}), - )?; - } - } - ToolProjectionEvent::Text(_) => {} - } - } - Ok(()) -} - -fn anthropic_stream_start( - stream: &mut impl Write, - request: &ResponseOptions, - id: &str, - prompt_tokens: u32, - cached_tokens: u32, -) -> Result<(), (u16, String)> { - let cached = cached_tokens.min(prompt_tokens); - let written = prompt_tokens - cached; - send_named_sse( - stream, - "message_start", - &json!({"type": "message_start", "message": {"id": id, "type": "message", "role": "assistant", "model": request.model_id, "content": [], "stop_reason": Value::Null, "stop_sequence": Value::Null, "usage": {"input_tokens": prompt_tokens - cached - written, "output_tokens": 0, "cache_read_input_tokens": cached, "cache_creation_input_tokens": written}}}), - ) -} - -fn responses_stream_response( - stream: &mut TcpStream, - state: &State, - request: ResponseOptions, - active: crate::runtime::ActiveGeneration, - id: &str, -) -> Result<(), (u16, String)> { - send_sse_headers(stream).map_err(|error| (500, error))?; - let created = unix_time(); - let message_id = random_id("msg_"); - let reasoning_id = random_id("rs_"); - let mut sequence = 0; - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.created", "response": {"id": id, "object": "response", "created_at": created, "status": "in_progress", "model": request.model_id, "output": []}}), - )?; - let mut reasoning_open = false; - let mut message_open = false; - let mut reasoning = String::new(); - let mut content = String::new(); - while let Ok(event) = active.events.recv() { - match event { - GenerationEvent::Chunk { - reasoning: true, - content: chunk, - } => { - if !request.reasoning_summary { - continue; - } - if !reasoning_open { - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.output_item.added", "output_index": 0, "item": {"id": reasoning_id, "type": "reasoning", "status": "in_progress", "summary": []}}), - )?; - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.reasoning_summary_part.added", "item_id": reasoning_id, "output_index": 0, "summary_index": 0, "part": {"type": "summary_text", "text": ""}}), - )?; - reasoning_open = true; - } - reasoning.push_str(&chunk); - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.reasoning_summary_text.delta", "item_id": reasoning_id, "output_index": 0, "summary_index": 0, "delta": chunk}), - )?; - } - GenerationEvent::Chunk { - reasoning: false, - content: chunk, - } => { - if request.has_tools { - content.push_str(&chunk); - continue; - } - if !message_open { - let output_index = usize::from(reasoning_open); - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.output_item.added", "output_index": output_index, "item": {"id": message_id, "type": "message", "status": "in_progress", "role": "assistant", "content": []}}), - )?; - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.content_part.added", "item_id": message_id, "output_index": output_index, "content_index": 0, "part": {"type": "output_text", "text": "", "annotations": []}}), - )?; - message_open = true; - } - content.push_str(&chunk); - let output_index = usize::from(reasoning_open); - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.output_text.delta", "item_id": message_id, "output_index": output_index, "content_index": 0, "delta": chunk}), - )?; - } - GenerationEvent::Finished(Ok(output)) => { - let (parsed_content, calls) = - parse_generated_tools(state, &output.message.content, Protocol::Responses); - if request.has_tools { - content = parsed_content; - } - let finish = if calls.is_empty() { - output.finish_reason - } else { - "tool_calls" - }; - let status = if finish == "length" { - "incomplete" - } else if finish == "error" { - "failed" - } else { - "completed" - }; - let mut terminal_items = Vec::new(); - let mut output_index = 0; - if reasoning_open { - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.reasoning_summary_text.done", "item_id": reasoning_id, "output_index": output_index, "summary_index": 0, "text": reasoning}), - )?; - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.reasoning_summary_part.done", "item_id": reasoning_id, "output_index": output_index, "summary_index": 0, "part": {"type": "summary_text", "text": reasoning}}), - )?; - let item = json!({"id": reasoning_id, "type": "reasoning", "status": status, "summary": [{"type": "summary_text", "text": reasoning}]}); - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.output_item.done", "output_index": output_index, "item": item}), - )?; - terminal_items.push(item); - output_index += 1; - } - if !content.is_empty() { - if !message_open { - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.output_item.added", "output_index": output_index, "item": {"id": message_id, "type": "message", "status": "in_progress", "role": "assistant", "content": []}}), - )?; - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.content_part.added", "item_id": message_id, "output_index": output_index, "content_index": 0, "part": {"type": "output_text", "text": "", "annotations": []}}), - )?; - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.output_text.delta", "item_id": message_id, "output_index": output_index, "content_index": 0, "delta": content}), - )?; - } - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.output_text.done", "item_id": message_id, "output_index": output_index, "content_index": 0, "text": content}), - )?; - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.content_part.done", "item_id": message_id, "output_index": output_index, "content_index": 0, "part": {"type": "output_text", "text": content, "annotations": []}}), - )?; - let item = json!({"id": message_id, "type": "message", "status": status, "role": "assistant", "content": [{"type": "output_text", "text": content, "annotations": []}]}); - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.output_item.done", "output_index": output_index, "item": item}), - )?; - terminal_items.push(item); - output_index += 1; - } - for call in &calls { - let item_id = random_id("fc_"); - let mut item = json!({"id": item_id, "type": "function_call", "status": status, "name": call.function.name, "call_id": call.id, "arguments": call.function.arguments}); - let mut added = item.clone(); - added["status"] = Value::String("in_progress".into()); - added["arguments"] = Value::String(String::new()); - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.output_item.added", "output_index": output_index, "item": added}), - )?; - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.function_call_arguments.delta", "item_id": item_id, "output_index": output_index, "delta": call.function.arguments}), - )?; - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.function_call_arguments.done", "item_id": item_id, "output_index": output_index, "name": call.function.name, "arguments": call.function.arguments}), - )?; - item["id"] = Value::String(item_id); - send_responses_sse( - stream, - &mut sequence, - json!({"type": "response.output_item.done", "output_index": output_index, "item": item}), - )?; - terminal_items.push(item); - output_index += 1; - } - let event_type = if finish == "length" { - "response.incomplete" - } else if finish == "error" { - "response.failed" - } else { - "response.completed" - }; - let cached = output.cached_tokens.min(output.prompt_tokens); - return send_responses_sse( - stream, - &mut sequence, - json!({"type": event_type, "response": {"id": id, "object": "response", "created_at": created, "status": status, "model": request.model_id, "output": terminal_items, "usage": {"input_tokens": output.prompt_tokens, "input_tokens_details": {"cached_tokens": cached, "cache_write_tokens": output.prompt_tokens - cached}, "output_tokens": output.completion_tokens, "output_tokens_details": {"reasoning_tokens": 0}, "total_tokens": output.prompt_tokens + output.completion_tokens}}}), - ); - } - GenerationEvent::Finished(Err(error)) => { - let _ = send_sse_error(stream, &error); - return Ok(()); - } - _ => {} - } - } - Err((500, "The model runtime stopped unexpectedly.".into())) -} - -#[allow(clippy::too_many_arguments)] -fn send_named_sse( - stream: &mut impl Write, - event: &str, - value: &Value, -) -> Result<(), (u16, String)> { - let body = serde_json::to_vec(value).map_err(|error| (500, error.to_string()))?; - stream - .write_all(b"event: ") - .and_then(|()| stream.write_all(event.as_bytes())) - .and_then(|()| stream.write_all(b"\ndata: ")) - .and_then(|()| stream.write_all(&body)) - .and_then(|()| stream.write_all(b"\n\n")) - .map_err(|error| (500, error.to_string())) -} - -fn send_responses_sse( - stream: &mut impl Write, - sequence: &mut u32, - value: Value, -) -> Result<(), (u16, String)> { - let mut object = value - .as_object() - .cloned() - .ok_or_else(|| (500, "Responses event is not an object".to_owned()))?; - let event_type = object.shift_remove("type").unwrap_or(Value::Null); - let mut ordered = Map::new(); - ordered.insert("type".into(), event_type); - ordered.insert("sequence_number".into(), Value::from(*sequence)); - ordered.extend(object); - *sequence += 1; - send_sse(stream, &Value::Object(ordered)).map_err(|error| (500, error)) -} - -fn stream_response_with_keepalive( - stream: &mut TcpStream, - state: &State, - request: ResponseOptions, - active: crate::runtime::ActiveGeneration, - id: &str, - keepalive_interval: Duration, -) -> Result<(), (u16, String)> { - let mut projector = ToolProjector::new(); - let mut output = None; - let mut prefilling = true; - let mut headers_sent = false; - let mut role_sent = false; - let mut last_keepalive = Instant::now(); - loop { - let event = match receive_stream_event( - stream, - &active, - prefilling && headers_sent, - &mut last_keepalive, - keepalive_interval, - ) { - Ok(Some(event)) => event, - Ok(None) => break, - Err(_) => return Ok(()), - }; - match event { - GenerationEvent::Loading => {} - GenerationEvent::Context { - tokens_per_second, .. - } => { - prefilling = tokens_per_second.is_none(); - if prefilling && !headers_sent { - send_sse_headers(stream).map_err(|error| (500, error))?; - headers_sent = true; - last_keepalive = Instant::now(); - } else if !prefilling { - send_stream_start(stream, &request, id, &mut headers_sent, &mut role_sent)?; - } - } - GenerationEvent::Chunk { reasoning, content } => { - prefilling = false; - send_stream_start(stream, &request, id, &mut headers_sent, &mut role_sent)?; - if request.has_tools && !reasoning { - let events = projector.push(&content, false, "call_"); - if send_chat_projection_events(stream, &request, id, events).is_err() { - active.cancel.store(true, Ordering::Relaxed); - return Ok(()); - } - if TOOL_SYNTAXES - .iter() - .any(|syntax| projector.raw.contains(syntax.tool_end)) - { - active.cancel.store(true, Ordering::Relaxed); - } - } else if !content.is_empty() { - let field = if reasoning { - "reasoning_content" - } else { - "content" - }; - let chunk = chunk_json(id, &request.model_id, json!({field: content}), None); - if send_sse(stream, &chunk).is_err() { - active.cancel.store(true, Ordering::Relaxed); - return Ok(()); - } - } - } - GenerationEvent::Finished(result) => { - match result { - Ok(result) => { - send_stream_start(stream, &request, id, &mut headers_sent, &mut role_sent)?; - output = Some(result); - } - Err(error) => { - if headers_sent { - let _ = send_sse_error(stream, &error); - return Ok(()); - } - return Err((500, error)); - } - } - break; - } - } - } - let output = match output { - Some(output) => output, - None if headers_sent => { - let _ = send_sse_error(stream, "The model runtime stopped unexpectedly."); - return Ok(()); - } - None => return Err((500, "The model runtime stopped unexpectedly.".into())), - }; - if request.has_tools { - let events = projector.push("", true, "call_"); - send_chat_projection_events(stream, &request, id, events)?; - } - let (_, calls) = parse_generated_tools_with_ids( - state, - &output.message.content, - Protocol::Chat, - &projector.ids, - ); - if request.has_tools && projector.ids.is_empty() && !calls.is_empty() { - send_sse( - stream, - &chunk_json( - id, - &request.model_id, - json!({"tool_calls": tool_calls_json(&calls)}), - None, - ), - ) - .map_err(|error| (500, error))?; - } - let finish = if calls.is_empty() { - output.finish_reason - } else { - "tool_calls" - }; - send_sse( - stream, - &chunk_json(id, &request.model_id, json!({}), Some(finish)), - ) - .map_err(|error| (500, error))?; - if request.include_usage { - let usage = json!({ - "id": id, - "object": "chat.completion.chunk", - "created": unix_time(), - "model": request.model_id, - "choices": [], - "usage": usage_json(output.prompt_tokens, output.cached_tokens, output.completion_tokens), - }); - send_sse(stream, &usage).map_err(|error| (500, error))?; - } - stream - .write_all(b"data: [DONE]\n\n") - .map_err(|error| (500, error.to_string())) -} - -fn send_chat_projection_events( - stream: &mut impl Write, - request: &ResponseOptions, - response_id: &str, - events: Vec, -) -> Result<(), (u16, String)> { - for event in events { - let delta = match event { - ToolProjectionEvent::Text(content) if !content.is_empty() => { - json!({"content": content}) - } - ToolProjectionEvent::Start { index, id, name } => json!({"tool_calls": [{ - "index": index, "id": id, "type": "function", - "function": {"name": name, "arguments": ""} - }]}), - ToolProjectionEvent::Arguments { index, fragment } => json!({ - "tool_calls": [{"index": index, "function": {"arguments": fragment}}] - }), - ToolProjectionEvent::End { .. } | ToolProjectionEvent::Text(_) => continue, - }; - send_sse( - stream, - &chunk_json(response_id, &request.model_id, delta, None), - ) - .map_err(|error| (500, error))?; - } - Ok(()) -} - -fn send_stream_start( - stream: &mut impl Write, - request: &ResponseOptions, - id: &str, - headers_sent: &mut bool, - role_sent: &mut bool, -) -> Result<(), (u16, String)> { - if !*headers_sent { - send_sse_headers(stream).map_err(|error| (500, error))?; - *headers_sent = true; - } - if !*role_sent { - let role = chunk_json(id, &request.model_id, json!({"role": "assistant"}), None); - send_sse(stream, &role).map_err(|error| (500, error))?; - *role_sent = true; - } - Ok(()) -} - -fn receive_stream_event( - stream: &mut impl Write, - active: &crate::runtime::ActiveGeneration, - prefilling: bool, - last_keepalive: &mut Instant, - keepalive_interval: Duration, -) -> Result, String> { - if !prefilling { - return Ok(active.events.recv().ok()); - } - loop { - if last_keepalive.elapsed() >= keepalive_interval { - if let Err(error) = stream.write_all(b": prefill\n\n") { - active.cancel.store(true, Ordering::Relaxed); - return Err(error.to_string()); - } - *last_keepalive = Instant::now(); - } - let remaining = keepalive_interval.saturating_sub(last_keepalive.elapsed()); - match active.events.recv_timeout(remaining) { - Ok(event) => return Ok(Some(event)), - Err(std::sync::mpsc::RecvTimeoutError::Timeout) => {} - Err(std::sync::mpsc::RecvTimeoutError::Disconnected) => return Ok(None), - } - } -} - -fn wait_for_output( - active: crate::runtime::ActiveGeneration, -) -> Result { - let mut content = String::new(); - while let Ok(event) = active.events.recv() { - match event { - GenerationEvent::Chunk { - reasoning: false, - content: chunk, - } => { - content.push_str(&chunk); - if TOOL_SYNTAXES - .iter() - .any(|syntax| content.contains(syntax.tool_end)) - { - active.cancel.store(true, Ordering::Relaxed); - } - } - GenerationEvent::Finished(result) => { - return result.map_err(|error| (500, error)); - } - _ => {} - } - } - Err((500, "The model runtime stopped unexpectedly.".into())) -} - -fn parse_generated_tools( - state: &State, - text: &str, - protocol: Protocol, -) -> (String, Vec) { - parse_generated_tools_with_ids(state, text, protocol, &[]) -} - -fn parse_generated_tools_with_ids( - state: &State, - text: &str, - protocol: Protocol, - streamed_ids: &[String], -) -> (String, Vec) { - let Some((start, syntax)) = TOOL_SYNTAXES - .iter() - .filter_map(|syntax| text.find(syntax.tool_start).map(|start| (start, *syntax))) - .min_by_key(|(start, _)| *start) - else { - return (text.to_owned(), Vec::new()); - }; - let Some(relative_end) = text[start..].find(syntax.tool_end) else { - return (text.to_owned(), Vec::new()); - }; - let end = start + relative_end + syntax.tool_end.len(); - let content = text[..start].trim_end(); - let raw = &text[content.len()..end]; - let mut calls = Vec::new(); - let mut cursor = raw.find(syntax.tool_start).unwrap() + syntax.tool_start.len(); - loop { - skip_text_whitespace(raw, &mut cursor); - if raw[cursor..].starts_with(syntax.tool_end) { - break; - } - if !raw[cursor..].starts_with(syntax.invoke_start) { - return (text.to_owned(), Vec::new()); - } - let Some(tag_end) = raw[cursor..].find('>').map(|end| cursor + end + 1) else { - return (text.to_owned(), Vec::new()); - }; - let Some(name) = dsml_attribute(&raw[cursor..tag_end], "name") else { - return (text.to_owned(), Vec::new()); - }; - cursor = tag_end; - let mut arguments = Map::new(); - loop { - skip_text_whitespace(raw, &mut cursor); - if raw[cursor..].starts_with(syntax.invoke_end) { - cursor += syntax.invoke_end.len(); - break; - } - let Some((name, value)) = parse_tool_parameter(raw, &mut cursor, syntax) else { - return (text.to_owned(), Vec::new()); - }; - arguments.insert(name, value); - } - calls.push(ApiToolCall { - id: String::new(), - function: ApiFunction { - name, - arguments: Value::Object(arguments).to_string(), - }, - }); - } - if calls.is_empty() { - return (text.to_owned(), Vec::new()); - } - let prefix = if protocol == Protocol::Anthropic { - "toolu_" - } else { - "call_" - }; - for (index, call) in calls.iter_mut().enumerate() { - call.id = streamed_ids - .get(index) - .cloned() - .unwrap_or_else(|| random_tool_id(prefix)); - } - if let Ok(mut memory) = state.tool_memory.lock() { - // ponytail: one process-local replay table; add LRU eviction if 100k live tool ids is measured insufficient. - if memory.len() >= 100_000 { - memory.clear(); - } - for call in &calls { - memory.insert(call.id.clone(), raw.to_owned()); - } - } - (content.to_owned(), calls) -} - -#[derive(Clone, Copy)] -struct ToolSyntax { - tool_start: &'static str, - tool_end: &'static str, - invoke_start: &'static str, - invoke_end: &'static str, - parameter_start: &'static str, - parameter_end: &'static str, -} - -const TOOL_SYNTAXES: [ToolSyntax; 3] = [ - ToolSyntax { - tool_start: "<|DSML|tool_calls>", - tool_end: "", - invoke_start: "<|DSML|invoke", - invoke_end: "", - parameter_start: "<|DSML|parameter", - parameter_end: "", - }, - ToolSyntax { - tool_start: "", - tool_end: "", - invoke_start: "", - parameter_start: "", - }, - ToolSyntax { - tool_start: "", - tool_end: "", - invoke_start: "", - parameter_start: "", - }, -]; - -fn parse_tool_parameter( - text: &str, - cursor: &mut usize, - syntax: ToolSyntax, -) -> Option<(String, Value)> { - if !text[*cursor..].starts_with(syntax.parameter_start) { - return None; - } - let tag_end = text[*cursor..].find('>').map(|end| *cursor + end + 1)?; - let tag = &text[*cursor..tag_end]; - let name = dsml_attribute(tag, "name")?; - let is_string = dsml_attribute(tag, "string"); - *cursor = tag_end; - let mut nested_start = *cursor; - skip_text_whitespace(text, &mut nested_start); - if is_string.is_none() && text[nested_start..].starts_with(syntax.parameter_start) { - *cursor = nested_start; - let mut nested = Map::new(); - loop { - skip_text_whitespace(text, cursor); - if !text[*cursor..].starts_with(syntax.parameter_start) { - break; - } - let (name, value) = parse_tool_parameter(text, cursor, syntax)?; - nested.insert(name, value); - } - skip_text_whitespace(text, cursor); - if !text[*cursor..].starts_with(syntax.parameter_end) { - return None; - } - *cursor += syntax.parameter_end.len(); - return Some((name, Value::Object(nested))); - } - let value_end = text[*cursor..] - .find(syntax.parameter_end) - .map(|end| *cursor + end)?; - let raw = &text[*cursor..value_end]; - *cursor = value_end + syntax.parameter_end.len(); - let value = if is_string.as_deref().unwrap_or("true") == "true" { - Value::String(unescape_dsml(raw)) - } else { - serde_json::from_str(raw).unwrap_or(Value::Null) - }; - Some((name, value)) -} - -fn skip_text_whitespace(text: &str, cursor: &mut usize) { - while text[*cursor..] - .chars() - .next() - .is_some_and(char::is_whitespace) - { - *cursor += text[*cursor..].chars().next().unwrap().len_utf8(); - } -} - -fn tool_calls_json(calls: &[ApiToolCall]) -> Value { - Value::Array( - calls - .iter() - .map(|call| { - json!({ - "id": call.id, - "type": "function", - "function": { - "name": call.function.name, - "arguments": call.function.arguments, - } - }) - }) - .collect(), - ) -} - fn installed_endpoint_models(models_path: &std::path::Path) -> Vec { model::installed_models(models_path) .into_iter() @@ -2711,247 +506,6 @@ fn model_alias(id: &str) -> Option { } } -fn content_text(value: &Value) -> String { - match value { - Value::String(text) => text.clone(), - Value::Array(parts) => parts - .iter() - .filter_map(|part| match part { - Value::String(text) => Some(text.as_str()), - Value::Object(object) => object.get("text").and_then(Value::as_str), - _ => None, - }) - .collect(), - _ => String::new(), - } -} - -fn escape_attribute(text: &str) -> String { - text.replace('&', "&") - .replace('<', "<") - .replace('>', ">") - .replace('"', """) -} - -fn escape_parameter(text: &str) -> String { - text.replace("", "</|DSML|parameter>") -} - -fn escape_json_parameter(text: &str) -> String { - text.replace("", "\\u003c/|DSML|parameter>") -} - -fn escape_tool_result(text: &str) -> String { - text.replace("", "</tool_result>") -} - -fn unescape_dsml(text: &str) -> String { - text.replace(""", "\"") - .replace(">", ">") - .replace("<", "<") - .replace("&", "&") -} - -fn usage_json(prompt: u32, cached: u32, completion: u32) -> Value { - let cached = cached.min(prompt); - json!({ - "prompt_tokens": prompt, - "completion_tokens": completion, - "total_tokens": prompt + completion, - "prompt_tokens_details": { - "cached_tokens": cached, - "cache_write_tokens": prompt - cached - } - }) -} - -fn chunk_json(id: &str, model: &str, delta: Value, finish: Option<&str>) -> Value { - json!({ - "id": id, - "object": "chat.completion.chunk", - "created": unix_time(), - "model": model, - "choices": [{"index": 0, "delta": delta, "finish_reason": finish}] - }) -} - -fn send_sse_headers(stream: &mut impl Write) -> Result<(), String> { - stream - .write_all( - b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nCache-Control: no-cache\r\nConnection: close\r\n\r\n", - ) - .map_err(|error| error.to_string()) -} - -fn send_sse(stream: &mut impl Write, value: &Value) -> Result<(), String> { - let body = serde_json::to_vec(value).map_err(|error| error.to_string())?; - stream - .write_all(b"data: ") - .map_err(|error| error.to_string())?; - stream.write_all(&body).map_err(|error| error.to_string())?; - stream.write_all(b"\n\n").map_err(|error| error.to_string()) -} - -fn send_sse_error(stream: &mut impl Write, message: &str) -> Result<(), String> { - let body = serde_json::to_vec(&json!({"error": {"message": message, "type": "server_error"}})) - .map_err(|error| error.to_string())?; - stream - .write_all(b"event: error\ndata: ") - .map_err(|error| error.to_string())?; - stream.write_all(&body).map_err(|error| error.to_string())?; - stream.write_all(b"\n\n").map_err(|error| error.to_string()) -} - -fn send_json(stream: &mut TcpStream, code: u16, value: &Value) -> Result<(), String> { - let mut body = serde_json::to_vec(value).map_err(|error| error.to_string())?; - body.push(b'\n'); - send_response(stream, code, Some("application/json"), &body) -} - -fn send_error(stream: &mut TcpStream, code: u16, message: &str) -> Result<(), String> { - send_json( - stream, - code, - &json!({"error": {"message": message, "type": "invalid_request_error"}}), - ) -} - -fn send_response( - stream: &mut TcpStream, - code: u16, - content_type: Option<&str>, - body: &[u8], -) -> Result<(), String> { - let reason = match code { - 200 => "OK", - 204 => "No Content", - 400 => "Bad Request", - 404 => "Not Found", - 409 => "Conflict", - 500 => "Internal Server Error", - _ => "Error", - }; - let mut header = format!( - "HTTP/1.1 {code} {reason}\r\nContent-Length: {}\r\n", - body.len() - ); - if let Some(content_type) = content_type { - header.push_str("Content-Type: "); - header.push_str(content_type); - header.push_str("\r\n"); - } - header.push_str("Connection: close\r\n\r\n"); - stream - .write_all(header.as_bytes()) - .and_then(|()| stream.write_all(body)) - .map_err(|error| error.to_string()) -} - -fn read_request(stream: &mut TcpStream) -> Result { - let mut bytes = Vec::new(); - let header_end = loop { - if bytes.len() >= MAX_HEADER_BYTES { - return Err("HTTP headers are too large".into()); - } - let mut chunk = [0_u8; 4096]; - let read = stream.read(&mut chunk).map_err(|error| error.to_string())?; - if read == 0 { - return Err("bad HTTP request".into()); - } - bytes.extend_from_slice(&chunk[..read]); - if let Some(end) = find_header_end(&bytes) { - break end; - } - }; - let header = std::str::from_utf8(&bytes[..header_end]) - .map_err(|_| "HTTP headers are not UTF-8".to_owned())?; - let mut lines = header.lines(); - let request_line = lines.next().ok_or_else(|| "bad HTTP request".to_owned())?; - let mut parts = request_line.split_whitespace(); - let method = parts - .next() - .ok_or_else(|| "bad HTTP request".to_owned())? - .to_owned(); - let path = parts - .next() - .ok_or_else(|| "bad HTTP request".to_owned())? - .to_owned(); - let length = lines - .find_map(|line| { - line.split_once(':').and_then(|(name, value)| { - name.eq_ignore_ascii_case("content-length") - .then(|| value.trim().parse::().ok()) - .flatten() - }) - }) - .unwrap_or(0); - if length > MAX_BODY_BYTES { - return Err("HTTP body is too large".into()); - } - while bytes.len() < header_end + length { - let mut chunk = [0_u8; 8192]; - let read = stream.read(&mut chunk).map_err(|error| error.to_string())?; - if read == 0 { - return Err("incomplete HTTP body".into()); - } - bytes.extend_from_slice(&chunk[..read]); - } - Ok(HttpRequest { - method, - path: path.split('?').next().unwrap_or(&path).to_owned(), - body: bytes[header_end..header_end + length].to_vec(), - }) -} - -fn find_header_end(bytes: &[u8]) -> Option { - bytes - .windows(4) - .position(|window| window == b"\r\n\r\n") - .map(|position| position + 4) - .or_else(|| { - bytes - .windows(2) - .position(|window| window == b"\n\n") - .map(|position| position + 2) - }) -} - -fn unix_time() -> u64 { - SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_secs() -} - -fn random_id(prefix: &str) -> String { - random_hex_id::<12>(prefix) -} - -fn random_tool_id(prefix: &str) -> String { - random_hex_id::<16>(prefix) -} - -fn random_hex_id(prefix: &str) -> String { - let mut bytes = [0_u8; N]; - if File::open("/dev/urandom") - .and_then(|mut file| file.read_exact(&mut bytes)) - .is_err() - { - let fallback = SystemTime::now() - .duration_since(UNIX_EPOCH) - .unwrap_or_default() - .as_nanos(); - bytes.copy_from_slice(&fallback.to_le_bytes()[..N]); - } - let mut id = String::with_capacity(prefix.len() + N * 2); - id.push_str(prefix); - for byte in bytes { - use std::fmt::Write as _; - let _ = write!(id, "{byte:02x}"); - } - id -} - fn default_role() -> String { "user".into() } diff --git a/src/server/http.rs b/src/server/http.rs new file mode 100644 index 0000000..0fa630b --- /dev/null +++ b/src/server/http.rs @@ -0,0 +1,178 @@ +use super::*; + +pub(super) fn send_sse_headers(stream: &mut impl Write) -> Result<(), String> { + stream + .write_all( + b"HTTP/1.1 200 OK\r\nContent-Type: text/event-stream\r\nCache-Control: no-cache\r\nConnection: close\r\n\r\n", + ) + .map_err(|error| error.to_string()) +} + +pub(super) fn send_sse(stream: &mut impl Write, value: &Value) -> Result<(), String> { + let body = serde_json::to_vec(value).map_err(|error| error.to_string())?; + stream + .write_all(b"data: ") + .map_err(|error| error.to_string())?; + stream.write_all(&body).map_err(|error| error.to_string())?; + stream.write_all(b"\n\n").map_err(|error| error.to_string()) +} + +pub(super) fn send_sse_error(stream: &mut impl Write, message: &str) -> Result<(), String> { + let body = serde_json::to_vec(&json!({"error": {"message": message, "type": "server_error"}})) + .map_err(|error| error.to_string())?; + stream + .write_all(b"event: error\ndata: ") + .map_err(|error| error.to_string())?; + stream.write_all(&body).map_err(|error| error.to_string())?; + stream.write_all(b"\n\n").map_err(|error| error.to_string()) +} + +pub(super) fn send_json(stream: &mut TcpStream, code: u16, value: &Value) -> Result<(), String> { + let mut body = serde_json::to_vec(value).map_err(|error| error.to_string())?; + body.push(b'\n'); + send_response(stream, code, Some("application/json"), &body) +} + +pub(super) fn send_error(stream: &mut TcpStream, code: u16, message: &str) -> Result<(), String> { + send_json( + stream, + code, + &json!({"error": {"message": message, "type": "invalid_request_error"}}), + ) +} + +pub(super) fn send_response( + stream: &mut TcpStream, + code: u16, + content_type: Option<&str>, + body: &[u8], +) -> Result<(), String> { + let reason = match code { + 200 => "OK", + 204 => "No Content", + 400 => "Bad Request", + 404 => "Not Found", + 409 => "Conflict", + 500 => "Internal Server Error", + _ => "Error", + }; + let mut header = format!( + "HTTP/1.1 {code} {reason}\r\nContent-Length: {}\r\n", + body.len() + ); + if let Some(content_type) = content_type { + header.push_str("Content-Type: "); + header.push_str(content_type); + header.push_str("\r\n"); + } + header.push_str("Connection: close\r\n\r\n"); + stream + .write_all(header.as_bytes()) + .and_then(|()| stream.write_all(body)) + .map_err(|error| error.to_string()) +} + +pub(super) fn read_request(stream: &mut TcpStream) -> Result { + let mut bytes = Vec::new(); + let header_end = loop { + if bytes.len() >= MAX_HEADER_BYTES { + return Err("HTTP headers are too large".into()); + } + let mut chunk = [0_u8; 4096]; + let read = stream.read(&mut chunk).map_err(|error| error.to_string())?; + if read == 0 { + return Err("bad HTTP request".into()); + } + bytes.extend_from_slice(&chunk[..read]); + if let Some(end) = find_header_end(&bytes) { + break end; + } + }; + let header = std::str::from_utf8(&bytes[..header_end]) + .map_err(|_| "HTTP headers are not UTF-8".to_owned())?; + let mut lines = header.lines(); + let request_line = lines.next().ok_or_else(|| "bad HTTP request".to_owned())?; + let mut parts = request_line.split_whitespace(); + let method = parts + .next() + .ok_or_else(|| "bad HTTP request".to_owned())? + .to_owned(); + let path = parts + .next() + .ok_or_else(|| "bad HTTP request".to_owned())? + .to_owned(); + let length = lines + .find_map(|line| { + line.split_once(':').and_then(|(name, value)| { + name.eq_ignore_ascii_case("content-length") + .then(|| value.trim().parse::().ok()) + .flatten() + }) + }) + .unwrap_or(0); + if length > MAX_BODY_BYTES { + return Err("HTTP body is too large".into()); + } + while bytes.len() < header_end + length { + let mut chunk = [0_u8; 8192]; + let read = stream.read(&mut chunk).map_err(|error| error.to_string())?; + if read == 0 { + return Err("incomplete HTTP body".into()); + } + bytes.extend_from_slice(&chunk[..read]); + } + Ok(HttpRequest { + method, + path: path.split('?').next().unwrap_or(&path).to_owned(), + body: bytes[header_end..header_end + length].to_vec(), + }) +} + +fn find_header_end(bytes: &[u8]) -> Option { + bytes + .windows(4) + .position(|window| window == b"\r\n\r\n") + .map(|position| position + 4) + .or_else(|| { + bytes + .windows(2) + .position(|window| window == b"\n\n") + .map(|position| position + 2) + }) +} + +pub(super) fn unix_time() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_secs() +} + +pub(super) fn random_id(prefix: &str) -> String { + random_hex_id::<12>(prefix) +} + +pub(super) fn random_tool_id(prefix: &str) -> String { + random_hex_id::<16>(prefix) +} + +fn random_hex_id(prefix: &str) -> String { + let mut bytes = [0_u8; N]; + if File::open("/dev/urandom") + .and_then(|mut file| file.read_exact(&mut bytes)) + .is_err() + { + let fallback = SystemTime::now() + .duration_since(UNIX_EPOCH) + .unwrap_or_default() + .as_nanos(); + bytes.copy_from_slice(&fallback.to_le_bytes()[..N]); + } + let mut id = String::with_capacity(prefix.len() + N * 2); + id.push_str(prefix); + for byte in bytes { + use std::fmt::Write as _; + let _ = write!(id, "{byte:02x}"); + } + id +} diff --git a/src/server/request.rs b/src/server/request.rs new file mode 100644 index 0000000..4a14746 --- /dev/null +++ b/src/server/request.rs @@ -0,0 +1,595 @@ +use super::*; + +pub(super) fn completion_request(mut value: Value) -> Result { + let object = value + .as_object_mut() + .ok_or_else(|| (400, "invalid JSON request".to_owned()))?; + let prompt = object + .remove("prompt") + .ok_or_else(|| (400, "missing prompt".to_owned()))?; + let prompt = match prompt { + Value::String(text) => text, + Value::Array(values) => values + .into_iter() + .next() + .and_then(|value| value.as_str().map(str::to_owned)) + .unwrap_or_default(), + _ => String::new(), + }; + object.insert( + "messages".into(), + json!([ + {"role": "system", "content": "You are a helpful assistant"}, + {"role": "user", "content": prompt} + ]), + ); + decode_chat_request(value) +} + +pub(super) fn anthropic_request(value: Value) -> Result { + let object = value + .as_object() + .ok_or_else(|| (400, "invalid JSON request".to_owned()))?; + let source = object + .get("messages") + .and_then(Value::as_array) + .ok_or_else(|| (400, "missing messages".to_owned()))?; + let mut messages = Vec::new(); + let system = object.get("system").map(content_text).unwrap_or_default(); + if !system.is_empty() { + messages.push(json!({"role": "system", "content": system})); + } + for message in source { + let role = message + .get("role") + .and_then(Value::as_str) + .unwrap_or("user"); + let Some(blocks) = message.get("content").and_then(Value::as_array) else { + messages.push(json!({"role": role, "content": message.get("content").cloned().unwrap_or(Value::Null)})); + continue; + }; + let mut text = String::new(); + let mut reasoning = String::new(); + let mut calls = Vec::new(); + for block in blocks { + match block.get("type").and_then(Value::as_str).unwrap_or("text") { + "text" => text.push_str(block.get("text").and_then(Value::as_str).unwrap_or("")), + "thinking" | "redacted_thinking" => reasoning.push_str( + block + .get("thinking") + .or_else(|| block.get("data")) + .and_then(Value::as_str) + .unwrap_or(""), + ), + "tool_use" => calls.push(json!({ + "id": block.get("id").and_then(Value::as_str).unwrap_or(""), + "type": "function", + "function": { + "name": block.get("name").and_then(Value::as_str).unwrap_or(""), + "arguments": block.get("input").cloned().unwrap_or_else(|| json!({})).to_string() + } + })), + "tool_result" => { + if !text.is_empty() { + messages.push(json!({"role": role, "content": std::mem::take(&mut text)})); + } + messages.push(json!({ + "role": "tool", + "content": block.get("content").map(content_text).unwrap_or_default(), + "tool_call_id": block.get("tool_use_id").and_then(Value::as_str).unwrap_or("") + })); + } + _ => {} + } + } + if !text.is_empty() || !reasoning.is_empty() || !calls.is_empty() { + messages.push(json!({ + "role": role, + "content": text, + "reasoning_content": reasoning, + "tool_calls": calls + })); + } + } + let mut request = Map::new(); + copy_request_fields(object, &mut request); + request.insert("messages".into(), Value::Array(messages)); + if let Some(stops) = object.get("stop_sequences") { + request.insert("stop".into(), stops.clone()); + } + if let Some(effort) = object + .get("output_config") + .and_then(|value| value.get("effort")) + { + request.insert("reasoning_effort".into(), effort.clone()); + } + request.insert( + "tools".into(), + normalize_anthropic_tools(object.get("tools")), + ); + decode_chat_request(Value::Object(request)) +} + +pub(super) fn responses_request(value: Value) -> Result { + let object = value + .as_object() + .ok_or_else(|| (400, "invalid JSON request".to_owned()))?; + for key in ["previous_response_id", "conversation"] { + if object.get(key).is_some_and(|value| !value.is_null()) { + return Err(( + 400, + format!("{key} is not supported; replay full input instead"), + )); + } + } + if let Some(choice) = object.get("tool_choice") { + match choice { + Value::String(choice) if choice == "none" || choice == "auto" => {} + Value::String(choice) => { + return Err((400, format!("tool_choice={choice} not supported"))); + } + Value::Object(_) => return Err((400, "forced tool_choice not supported".into())), + _ => {} + } + } + let input = object + .get("input") + .ok_or_else(|| (400, "missing input".to_owned()))?; + let mut messages = Vec::new(); + match input { + Value::String(text) => messages.push(json!({"role": "user", "content": text})), + Value::Array(items) => { + let mut pending_reasoning = String::new(); + for item in items { + let item = item + .as_object() + .ok_or_else(|| (400, "invalid JSON request".to_owned()))?; + if item + .get("status") + .and_then(Value::as_str) + .is_some_and(|status| status != "completed") + { + return Err((400, "invalid JSON request".into())); + } + let kind = item + .get("type") + .and_then(Value::as_str) + .unwrap_or("message"); + let role = item.get("role").and_then(Value::as_str).unwrap_or("user"); + let consumes_reasoning = (kind == "message" && role == "assistant") + || matches!( + kind, + "function_call" + | "custom_tool_call" + | "local_shell_call" + | "web_search_call" + | "tool_search_call" + | "image_generation_call" + ); + let bookkeeping = matches!(kind, "compaction" | "context_compaction"); + if !consumes_reasoning && !bookkeeping && !pending_reasoning.is_empty() { + messages.push(reasoning_message(std::mem::take(&mut pending_reasoning))); + } + match kind { + "message" => { + let mut message = json!({ + "role": role, + "content": responses_content_text(item.get("content").unwrap_or(&Value::Null))? + }); + if role == "assistant" && !pending_reasoning.is_empty() { + message["reasoning_content"] = + Value::String(std::mem::take(&mut pending_reasoning)); + } + messages.push(message); + } + "function_call" | "custom_tool_call" => { + let name = if kind == "function_call" { + format!( + "{}{}", + item.get("namespace").and_then(Value::as_str).unwrap_or(""), + item.get("name").and_then(Value::as_str).unwrap_or("") + ) + } else { + item.get("name") + .and_then(Value::as_str) + .unwrap_or("") + .to_owned() + }; + push_responses_call( + &mut messages, + &mut pending_reasoning, + item, + &name, + item.get("arguments").or_else(|| item.get("input")), + ); + } + "function_call_output" | "custom_tool_call_output" => messages.push(json!({ + "role": "tool", + "content": responses_output_text(item.get("output"))?, + "tool_call_id": item.get("call_id").or_else(|| item.get("id")).and_then(Value::as_str).unwrap_or("") + })), + "reasoning" => { + for value in [item.get("summary"), item.get("content")] + .into_iter() + .flatten() + { + let text = responses_content_text(value)?; + if !pending_reasoning.is_empty() && !text.is_empty() { + pending_reasoning.push('\n'); + } + pending_reasoning.push_str(&text); + } + } + "local_shell_call" + | "web_search_call" + | "tool_search_call" + | "image_generation_call" => { + let name = match kind { + "local_shell_call" => "local_shell", + "tool_search_call" => "tool_search", + other => other, + }; + push_responses_call( + &mut messages, + &mut pending_reasoning, + item, + name, + item.get("action") + .or_else(|| item.get("arguments")) + .or_else(|| item.get("input")), + ); + } + "local_shell_call_output" + | "web_search_call_output" + | "tool_search_output" + | "tool_search_call_output" + | "image_generation_call_output" => { + let content = if let Some(value) = + item.get("output").or_else(|| item.get("result")) + { + responses_output_text(Some(value))? + } else { + item.get("tools").map(Value::to_string).unwrap_or_default() + }; + messages.push(json!({ + "role": "tool", + "content": content, + "tool_call_id": item.get("call_id").or_else(|| item.get("id")).and_then(Value::as_str).unwrap_or("") + })); + } + "compaction" | "context_compaction" => {} + _ => return Err((400, "invalid JSON request".into())), + } + } + if !pending_reasoning.is_empty() { + messages.push(reasoning_message(pending_reasoning)); + } + } + _ => return Err((400, "invalid JSON request".into())), + } + if let Some(instructions) = object.get("instructions") { + let instructions = match instructions { + Value::Null => String::new(), + Value::String(text) => text.clone(), + _ => return Err((400, "invalid JSON request".into())), + }; + if !instructions.is_empty() { + messages.insert(0, json!({"role": "system", "content": instructions})); + } + } + let mut request = Map::new(); + copy_request_fields(object, &mut request); + request.insert("messages".into(), Value::Array(messages)); + if let Some(tokens) = object.get("max_output_tokens") { + request.insert("max_tokens".into(), tokens.clone()); + } + if let Some(effort) = object + .get("reasoning") + .and_then(|value| value.get("effort")) + { + request.insert("reasoning_effort".into(), effort.clone()); + } + request.insert( + "tools".into(), + normalize_responses_tools(object.get("tools")), + ); + decode_chat_request(Value::Object(request)) +} + +fn reasoning_message(reasoning: String) -> Value { + json!({"role": "assistant", "content": "", "reasoning_content": reasoning}) +} + +fn push_responses_call( + messages: &mut Vec, + pending_reasoning: &mut String, + item: &Map, + name: &str, + arguments: Option<&Value>, +) { + let call = json!({ + "id": item.get("call_id").or_else(|| item.get("id")).and_then(Value::as_str).unwrap_or(""), + "type": "function", + "function": {"name": name, "arguments": raw_json_text(arguments)} + }); + if let Some(last) = messages.last_mut() + && last.get("role").and_then(Value::as_str) == Some("assistant") + { + if !pending_reasoning.is_empty() + && last + .get("reasoning_content") + .and_then(Value::as_str) + .is_none_or(str::is_empty) + { + last["reasoning_content"] = Value::String(std::mem::take(pending_reasoning)); + } + if !last.get("tool_calls").is_some_and(Value::is_array) { + last["tool_calls"] = json!([]); + } + last["tool_calls"].as_array_mut().unwrap().push(call); + } else { + messages.push(json!({ + "role": "assistant", + "content": "", + "reasoning_content": std::mem::take(pending_reasoning), + "tool_calls": [call] + })); + } +} + +fn raw_json_text(value: Option<&Value>) -> String { + match value { + Some(Value::String(text)) => text.clone(), + Some(value) => value.to_string(), + None => "{}".into(), + } +} + +fn responses_output_text(value: Option<&Value>) -> Result { + match value { + Some(Value::Array(_)) => responses_content_text(value.unwrap()), + Some(Value::String(text)) => Ok(text.clone()), + Some(value) => Ok(value.to_string()), + None => Ok(String::new()), + } +} + +fn responses_content_text(value: &Value) -> Result { + match value { + Value::String(text) => Ok(text.clone()), + Value::Null => Ok(String::new()), + Value::Array(parts) => { + let mut text = String::new(); + for part in parts { + match part { + Value::String(part) => text.push_str(part), + Value::Object(part) + if matches!( + part.get("type").and_then(Value::as_str), + Some( + "input_text" + | "output_text" + | "text" + | "summary_text" + | "reasoning_text" + ) + ) && matches!( + part.get("text"), + Some(Value::String(_) | Value::Null) + ) => + { + if let Some(value) = part.get("text").and_then(Value::as_str) { + text.push_str(value); + } + } + _ => return Err((400, "invalid JSON request".into())), + } + } + Ok(text) + } + _ => Err((400, "invalid JSON request".into())), + } +} + +fn copy_request_fields(source: &Map, target: &mut Map) { + for key in [ + "model", + "max_tokens", + "max_completion_tokens", + "temperature", + "top_p", + "min_p", + "top_k", + "seed", + "stream", + "stream_options", + "thinking", + "think", + "reasoning_effort", + "tool_choice", + "stop", + ] { + if let Some(value) = source.get(key) { + target.insert(key.into(), value.clone()); + } + } +} + +fn normalize_anthropic_tools(tools: Option<&Value>) -> Value { + Value::Array( + tools + .and_then(Value::as_array) + .into_iter() + .flatten() + .map(|tool| { + json!({"type": "function", "function": { + "name": tool.get("name").and_then(Value::as_str).unwrap_or(""), + "description": tool.get("description").and_then(Value::as_str).unwrap_or(""), + "parameters": tool.get("input_schema").cloned().unwrap_or_else(|| json!({"type": "object"})) + }}) + }) + .collect(), + ) +} + +fn normalize_responses_tools(tools: Option<&Value>) -> Value { + Value::Array( + tools + .and_then(Value::as_array) + .into_iter() + .flatten() + .filter(|tool| tool.get("type").and_then(Value::as_str) == Some("function")) + .map(|tool| { + json!({"type": "function", "function": { + "name": tool.get("name").and_then(Value::as_str).unwrap_or(""), + "description": tool.get("description").and_then(Value::as_str).unwrap_or(""), + "parameters": tool.get("parameters").cloned().unwrap_or_else(|| json!({"type": "object"})) + }}) + }) + .collect(), + ) +} + +pub(super) fn decode_chat_request(value: Value) -> Result { + serde_json::from_value(value).map_err(|_| (400, "invalid JSON request".to_owned())) +} + +pub(super) fn raw_tool_schemas( + body: &[u8], + protocol: Protocol, +) -> Result, (u16, String)> { + let request: RawToolsRequest<'_> = + serde_json::from_slice(body).map_err(|_| (400, "invalid JSON request".to_owned()))?; + request + .tools + .into_iter() + .map(|tool| { + if protocol == Protocol::Chat { + let wrapper: RawOpenAiTool<'_> = serde_json::from_str(tool.get()) + .map_err(|_| (400, "invalid JSON request".to_owned()))?; + Ok(wrapper.function.unwrap_or(tool).get().to_owned()) + } else { + Ok(tool.get().to_owned()) + } + }) + .collect() +} + +pub(super) fn parse_chat_request( + state: &State, + request: ChatRequest, + protocol: Protocol, +) -> Result { + if request.messages.is_empty() { + return Err((400, "missing messages".into())); + } + let requested_id = request.model.as_deref().unwrap_or_default(); + let preferences = state + .preferences + .read() + .map_err(|_| (500, "preferences are unavailable".to_owned()))? + .clone(); + let model = if requested_id.is_empty() { + ModelChoice::from_id(&preferences.selected_model) + } else { + model_alias(requested_id) + } + .ok_or_else(|| (400, format!("unknown model: {requested_id}")))?; + if !installed_endpoint_models(&state.models_path).contains(&model) { + return Err((400, format!("model is not installed and verified: {model}"))); + } + + let mut generation = preferences.generation().map_err(|error| (500, error))?; + generation.system_prompt.clear(); + generation.max_generated_tokens = request + .max_completion_tokens + .or(request.max_tokens) + .unwrap_or(generation.max_generated_tokens); + if generation.max_generated_tokens <= 0 { + return Err((400, "max_tokens must be positive".into())); + } + generation.temperature = request.temperature.or(Some(1.0)); + generation.top_p = request.top_p.or(Some(1.0)); + generation.min_p = request.min_p.or(Some(0.05)); + generation.seed = request.seed.filter(|seed| *seed > 0); + generation.reasoning_mode = request_reasoning(&request, requested_id)?; + + let tools_enabled = (!request.tools.is_empty() || !request.tool_schemas.is_empty()) + && request.tool_choice.as_ref().and_then(Value::as_str) != Some("none"); + let (system, messages) = render_messages( + state, + &request.messages, + &request.tools, + &request.tool_schemas, + tools_enabled, + protocol, + )?; + generation.system_prompt = system; + let runtime = preferences.runtime().map_err(|error| (500, error))?; + let mut effective = effective_settings(model, &generation, &runtime, &state.models_path) + .map_err(|error| (400, error))?; + effective.turn.top_k = request.top_k.unwrap_or(0); + if effective.turn.top_k < 0 { + return Err((400, "top_k must not be negative".into())); + } + effective.turn.stops = match request.stop { + Some(OneOrMany::One(stop)) => vec![stop], + Some(OneOrMany::Many(stops)) => stops, + None => Vec::new(), + }; + effective.turn.stops.retain(|stop| !stop.is_empty()); + + Ok(ParsedRequest { + protocol, + model_id: if requested_id.is_empty() { + model.id().to_owned() + } else { + requested_id.to_owned() + }, + messages, + turn: effective.turn, + engine: effective.engine, + idle_timeout: Duration::from_secs(preferences.idle_timeout_minutes.max(1) as u64 * 60), + stream: request.stream, + include_usage: request + .stream_options + .is_some_and(|options| options.include_usage), + has_tools: tools_enabled, + }) +} + +fn request_reasoning( + request: &ChatRequest, + model_id: &str, +) -> Result { + let explicit_thinking = request + .think + .or_else(|| request.thinking.as_ref().and_then(thinking_enabled)); + let mut reasoning = match request.reasoning_effort.as_deref() { + Some("max") => ReasoningMode::Max, + Some("none") => ReasoningMode::Direct, + Some("xhigh" | "high" | "medium" | "low" | "minimal") | None => ReasoningMode::High, + Some(value) => return Err((400, format!("unsupported reasoning_effort: {value}"))), + }; + if explicit_thinking == Some(false) + || (explicit_thinking.is_none() + && matches!( + model_id, + "deepseek-chat" | "glm-5.2-chat" | "glm-5.2-no-think" | "glm-5.2-nothink" + )) + { + reasoning = ReasoningMode::Direct; + } + Ok(reasoning) +} + +fn thinking_enabled(value: &Value) -> Option { + value.as_bool().or_else(|| { + value.as_str().map(|value| value != "disabled").or_else(|| { + value + .get("type") + .and_then(Value::as_str) + .map(|value| value != "disabled") + }) + }) +} diff --git a/src/server/response.rs b/src/server/response.rs new file mode 100644 index 0000000..47ca807 --- /dev/null +++ b/src/server/response.rs @@ -0,0 +1,974 @@ +use super::*; + +pub(super) fn final_response( + stream: &mut TcpStream, + state: &State, + request: ResponseOptions, + active: crate::runtime::ActiveGeneration, + id: &str, +) -> Result<(), (u16, String)> { + let output = wait_for_output(active)?; + let (content, calls) = parse_generated_tools(state, &output.message.content, request.protocol); + let finish = if calls.is_empty() { + output.finish_reason + } else { + "tool_calls" + }; + let reasoning = output.message.reasoning.filter(|value| !value.is_empty()); + let usage = usage_json( + output.prompt_tokens, + output.cached_tokens, + output.completion_tokens, + ); + let body = match request.protocol { + Protocol::Chat => { + let mut message = json!({"role": "assistant", "content": content}); + if let Some(reasoning) = reasoning { + message["reasoning_content"] = Value::String(reasoning); + } + if !calls.is_empty() { + message["tool_calls"] = tool_calls_json(&calls); + } + json!({ + "id": id, "object": "chat.completion", "created": unix_time(), + "model": request.model_id, + "choices": [{"index": 0, "message": message, "finish_reason": finish}], + "usage": usage, + }) + } + Protocol::Completion => json!({ + "id": id, "object": "text_completion", "created": unix_time(), + "model": request.model_id, + "choices": [{"text": content, "index": 0, "finish_reason": finish}], + "usage": usage, + }), + Protocol::Anthropic => { + let mut blocks = Vec::new(); + if let Some(reasoning) = reasoning { + blocks.push(json!({"type": "thinking", "thinking": reasoning, "signature": id})); + } + if !content.is_empty() { + blocks.push(json!({"type": "text", "text": content})); + } + for call in &calls { + blocks.push(json!({ + "type": "tool_use", "id": call.id, "name": call.function.name, + "input": serde_json::from_str::(&call.function.arguments).unwrap_or_else(|_| json!({})) + })); + } + if blocks.is_empty() || (blocks.iter().all(|block| block["type"] == "thinking")) { + blocks.push(json!({"type": "text", "text": ""})); + } + let cached = output.cached_tokens.min(output.prompt_tokens); + let written = output.prompt_tokens - cached; + json!({ + "id": id, "type": "message", "role": "assistant", "model": request.model_id, + "content": blocks, + "stop_reason": if finish == "tool_calls" { "tool_use" } else if finish == "length" { "max_tokens" } else { "end_turn" }, + "stop_sequence": Value::Null, + "usage": { + "input_tokens": output.prompt_tokens - cached - written, + "output_tokens": output.completion_tokens, + "cache_read_input_tokens": cached, + "cache_creation_input_tokens": written + } + }) + } + Protocol::Responses => { + let status = if finish == "length" { + "incomplete" + } else if finish == "error" { + "failed" + } else { + "completed" + }; + let mut items = Vec::new(); + if let Some(reasoning) = reasoning { + items.push(json!({ + "id": random_id("rs_"), "type": "reasoning", "status": status, + "summary": [{"type": "summary_text", "text": reasoning}] + })); + } + if !content.is_empty() { + items.push(json!({ + "id": random_id("msg_"), "type": "message", "status": status, + "role": "assistant", "content": [{"type": "output_text", "text": content, "annotations": []}] + })); + } + for call in &calls { + items.push(json!({ + "id": random_id("fc_"), "type": "function_call", "status": status, + "name": call.function.name, "call_id": call.id, "arguments": call.function.arguments + })); + } + json!({ + "id": id, "object": "response", "created_at": unix_time(), "status": status, + "model": request.model_id, "output": items, + "usage": { + "input_tokens": output.prompt_tokens, + "input_tokens_details": {"cached_tokens": output.cached_tokens.min(output.prompt_tokens), "cache_write_tokens": output.prompt_tokens - output.cached_tokens.min(output.prompt_tokens)}, + "output_tokens": output.completion_tokens, + "output_tokens_details": {"reasoning_tokens": 0}, + "total_tokens": output.prompt_tokens + output.completion_tokens + } + }) + } + }; + send_json(stream, 200, &body).map_err(|error| (500, error)) +} + +pub(super) fn stream_response( + stream: &mut TcpStream, + state: &State, + request: ResponseOptions, + active: crate::runtime::ActiveGeneration, + id: &str, +) -> Result<(), (u16, String)> { + if request.protocol != Protocol::Chat { + return structured_stream_response(stream, state, request, active, id); + } + stream_response_with_keepalive( + stream, + state, + request, + active, + id, + PREFILL_KEEPALIVE_INTERVAL, + ) +} + +fn structured_stream_response( + stream: &mut TcpStream, + state: &State, + request: ResponseOptions, + active: crate::runtime::ActiveGeneration, + id: &str, +) -> Result<(), (u16, String)> { + match request.protocol { + Protocol::Completion => completion_stream_response(stream, request, active, id), + Protocol::Anthropic => anthropic_stream_response(stream, state, request, active, id), + Protocol::Responses => responses_stream_response(stream, state, request, active, id), + Protocol::Chat => unreachable!(), + } +} + +fn completion_stream_response( + stream: &mut TcpStream, + request: ResponseOptions, + active: crate::runtime::ActiveGeneration, + id: &str, +) -> Result<(), (u16, String)> { + send_sse_headers(stream).map_err(|error| (500, error))?; + while let Ok(event) = active.events.recv() { + match event { + GenerationEvent::Chunk { + reasoning: false, + content, + } if !content.is_empty() => send_sse( + stream, + &json!({ + "id": id, "object": "text_completion", "created": unix_time(), + "model": request.model_id, + "choices": [{"text": content, "index": 0, "finish_reason": Value::Null}] + }), + ) + .map_err(|error| (500, error))?, + GenerationEvent::Finished(Ok(output)) => { + send_sse( + stream, + &json!({ + "id": id, "object": "text_completion", "created": unix_time(), + "model": request.model_id, + "choices": [{"text": "", "index": 0, "finish_reason": output.finish_reason}] + }), + ) + .map_err(|error| (500, error))?; + if request.include_usage { + send_sse( + stream, + &json!({ + "id": id, "object": "text_completion", "created": unix_time(), + "model": request.model_id, "choices": [], + "usage": usage_json(output.prompt_tokens, output.cached_tokens, output.completion_tokens) + }), + ) + .map_err(|error| (500, error))?; + } + return stream + .write_all(b"data: [DONE]\n\n") + .map_err(|error| (500, error.to_string())); + } + GenerationEvent::Finished(Err(error)) => { + let _ = send_sse_error(stream, &error); + return Ok(()); + } + _ => {} + } + } + Err((500, "The model runtime stopped unexpectedly.".into())) +} + +fn anthropic_stream_response( + stream: &mut TcpStream, + state: &State, + request: ResponseOptions, + active: crate::runtime::ActiveGeneration, + id: &str, +) -> Result<(), (u16, String)> { + send_sse_headers(stream).map_err(|error| (500, error))?; + let mut prompt_tokens = 0; + let mut started = false; + let mut block = None::<(usize, bool)>; + let mut next_index = 0; + let mut projector = ToolProjector::new(); + let mut tool_indices = Vec::new(); + while let Ok(event) = active.events.recv() { + match event { + GenerationEvent::Context { + used, + tokens_per_second, + .. + } => { + if tokens_per_second.is_none() { + prompt_tokens = used; + } else if !started { + anthropic_stream_start(stream, &request, id, prompt_tokens, 0)?; + started = true; + } + } + GenerationEvent::Chunk { reasoning, content } => { + if !started { + anthropic_stream_start(stream, &request, id, prompt_tokens, 0)?; + started = true; + } + if !reasoning && request.has_tools { + let events = projector.push(&content, false, "toolu_"); + send_anthropic_projection_events( + stream, + events, + &mut block, + &mut next_index, + &mut tool_indices, + )?; + continue; + } + if block.is_some_and(|(_, current_reasoning)| current_reasoning != reasoning) { + let (index, _) = block.take().unwrap(); + send_named_sse( + stream, + "content_block_stop", + &json!({"type": "content_block_stop", "index": index}), + )?; + } + let index = if let Some((index, _)) = block { + index + } else { + let index = next_index; + next_index += 1; + send_named_sse( + stream, + "content_block_start", + &json!({ + "type": "content_block_start", "index": index, + "content_block": if reasoning { json!({"type": "thinking", "thinking": "", "signature": ""}) } else { json!({"type": "text", "text": ""}) } + }), + )?; + block = Some((index, reasoning)); + index + }; + let delta = if reasoning { + json!({"type": "thinking_delta", "thinking": content}) + } else { + json!({"type": "text_delta", "text": content}) + }; + send_named_sse( + stream, + "content_block_delta", + &json!({"type": "content_block_delta", "index": index, "delta": delta}), + )?; + } + GenerationEvent::Finished(Ok(output)) => { + if !started { + anthropic_stream_start( + stream, + &request, + id, + output.prompt_tokens, + output.cached_tokens, + )?; + } + if request.has_tools { + let events = projector.push("", true, "toolu_"); + send_anthropic_projection_events( + stream, + events, + &mut block, + &mut next_index, + &mut tool_indices, + )?; + } + if let Some((index, _)) = block.take() { + send_named_sse( + stream, + "content_block_stop", + &json!({"type": "content_block_stop", "index": index}), + )?; + } + let (_, calls) = parse_generated_tools_with_ids( + state, + &output.message.content, + Protocol::Anthropic, + &projector.ids, + ); + if projector.ids.is_empty() { + for call in &calls { + send_named_sse( + stream, + "content_block_start", + &json!({"type": "content_block_start", "index": next_index, "content_block": {"type": "tool_use", "id": call.id, "name": call.function.name, "input": {}}}), + )?; + send_named_sse( + stream, + "content_block_delta", + &json!({"type": "content_block_delta", "index": next_index, "delta": {"type": "input_json_delta", "partial_json": call.function.arguments}}), + )?; + send_named_sse( + stream, + "content_block_stop", + &json!({"type": "content_block_stop", "index": next_index}), + )?; + next_index += 1; + } + } + let finish = if calls.is_empty() { + output.finish_reason + } else { + "tool_calls" + }; + send_named_sse( + stream, + "message_delta", + &json!({"type": "message_delta", "delta": {"stop_reason": if finish == "tool_calls" { "tool_use" } else if finish == "length" { "max_tokens" } else { "end_turn" }, "stop_sequence": Value::Null}, "usage": {"output_tokens": output.completion_tokens}}), + )?; + return send_named_sse(stream, "message_stop", &json!({"type": "message_stop"})); + } + GenerationEvent::Finished(Err(error)) => { + return send_named_sse( + stream, + "error", + &json!({"type": "error", "error": {"type": "api_error", "message": error}}), + ); + } + _ => {} + } + } + Err((500, "The model runtime stopped unexpectedly.".into())) +} + +fn send_anthropic_projection_events( + stream: &mut impl Write, + events: Vec, + block: &mut Option<(usize, bool)>, + next_index: &mut usize, + tool_indices: &mut Vec, +) -> Result<(), (u16, String)> { + for event in events { + match event { + ToolProjectionEvent::Text(text) if !text.is_empty() => { + if block.is_some_and(|(_, reasoning)| reasoning) { + let (index, _) = block.take().unwrap(); + send_named_sse( + stream, + "content_block_stop", + &json!({"type": "content_block_stop", "index": index}), + )?; + } + let index = if let Some((index, _)) = *block { + index + } else { + let index = *next_index; + *next_index += 1; + send_named_sse( + stream, + "content_block_start", + &json!({"type": "content_block_start", "index": index, "content_block": {"type": "text", "text": ""}}), + )?; + *block = Some((index, false)); + index + }; + send_named_sse( + stream, + "content_block_delta", + &json!({"type": "content_block_delta", "index": index, "delta": {"type": "text_delta", "text": text}}), + )?; + } + ToolProjectionEvent::Start { index, id, name } => { + if let Some((open_index, _)) = block.take() { + send_named_sse( + stream, + "content_block_stop", + &json!({"type": "content_block_stop", "index": open_index}), + )?; + } + let content_index = *next_index; + *next_index += 1; + if tool_indices.len() == index { + tool_indices.push(content_index); + } + send_named_sse( + stream, + "content_block_start", + &json!({"type": "content_block_start", "index": content_index, "content_block": {"type": "tool_use", "id": id, "name": name, "input": {}}}), + )?; + } + ToolProjectionEvent::Arguments { index, fragment } => { + if let Some(content_index) = tool_indices.get(index) { + send_named_sse( + stream, + "content_block_delta", + &json!({"type": "content_block_delta", "index": content_index, "delta": {"type": "input_json_delta", "partial_json": fragment}}), + )?; + } + } + ToolProjectionEvent::End { index } => { + if let Some(content_index) = tool_indices.get(index) { + send_named_sse( + stream, + "content_block_stop", + &json!({"type": "content_block_stop", "index": content_index}), + )?; + } + } + ToolProjectionEvent::Text(_) => {} + } + } + Ok(()) +} + +fn anthropic_stream_start( + stream: &mut impl Write, + request: &ResponseOptions, + id: &str, + prompt_tokens: u32, + cached_tokens: u32, +) -> Result<(), (u16, String)> { + let cached = cached_tokens.min(prompt_tokens); + let written = prompt_tokens - cached; + send_named_sse( + stream, + "message_start", + &json!({"type": "message_start", "message": {"id": id, "type": "message", "role": "assistant", "model": request.model_id, "content": [], "stop_reason": Value::Null, "stop_sequence": Value::Null, "usage": {"input_tokens": prompt_tokens - cached - written, "output_tokens": 0, "cache_read_input_tokens": cached, "cache_creation_input_tokens": written}}}), + ) +} + +fn responses_stream_response( + stream: &mut TcpStream, + state: &State, + request: ResponseOptions, + active: crate::runtime::ActiveGeneration, + id: &str, +) -> Result<(), (u16, String)> { + send_sse_headers(stream).map_err(|error| (500, error))?; + let created = unix_time(); + let message_id = random_id("msg_"); + let reasoning_id = random_id("rs_"); + let mut sequence = 0; + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.created", "response": {"id": id, "object": "response", "created_at": created, "status": "in_progress", "model": request.model_id, "output": []}}), + )?; + let mut reasoning_open = false; + let mut message_open = false; + let mut reasoning = String::new(); + let mut content = String::new(); + while let Ok(event) = active.events.recv() { + match event { + GenerationEvent::Chunk { + reasoning: true, + content: chunk, + } => { + if !request.reasoning_summary { + continue; + } + if !reasoning_open { + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.output_item.added", "output_index": 0, "item": {"id": reasoning_id, "type": "reasoning", "status": "in_progress", "summary": []}}), + )?; + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.reasoning_summary_part.added", "item_id": reasoning_id, "output_index": 0, "summary_index": 0, "part": {"type": "summary_text", "text": ""}}), + )?; + reasoning_open = true; + } + reasoning.push_str(&chunk); + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.reasoning_summary_text.delta", "item_id": reasoning_id, "output_index": 0, "summary_index": 0, "delta": chunk}), + )?; + } + GenerationEvent::Chunk { + reasoning: false, + content: chunk, + } => { + if request.has_tools { + content.push_str(&chunk); + continue; + } + if !message_open { + let output_index = usize::from(reasoning_open); + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.output_item.added", "output_index": output_index, "item": {"id": message_id, "type": "message", "status": "in_progress", "role": "assistant", "content": []}}), + )?; + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.content_part.added", "item_id": message_id, "output_index": output_index, "content_index": 0, "part": {"type": "output_text", "text": "", "annotations": []}}), + )?; + message_open = true; + } + content.push_str(&chunk); + let output_index = usize::from(reasoning_open); + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.output_text.delta", "item_id": message_id, "output_index": output_index, "content_index": 0, "delta": chunk}), + )?; + } + GenerationEvent::Finished(Ok(output)) => { + let (parsed_content, calls) = + parse_generated_tools(state, &output.message.content, Protocol::Responses); + if request.has_tools { + content = parsed_content; + } + let finish = if calls.is_empty() { + output.finish_reason + } else { + "tool_calls" + }; + let status = if finish == "length" { + "incomplete" + } else if finish == "error" { + "failed" + } else { + "completed" + }; + let mut terminal_items = Vec::new(); + let mut output_index = 0; + if reasoning_open { + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.reasoning_summary_text.done", "item_id": reasoning_id, "output_index": output_index, "summary_index": 0, "text": reasoning}), + )?; + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.reasoning_summary_part.done", "item_id": reasoning_id, "output_index": output_index, "summary_index": 0, "part": {"type": "summary_text", "text": reasoning}}), + )?; + let item = json!({"id": reasoning_id, "type": "reasoning", "status": status, "summary": [{"type": "summary_text", "text": reasoning}]}); + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.output_item.done", "output_index": output_index, "item": item}), + )?; + terminal_items.push(item); + output_index += 1; + } + if !content.is_empty() { + if !message_open { + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.output_item.added", "output_index": output_index, "item": {"id": message_id, "type": "message", "status": "in_progress", "role": "assistant", "content": []}}), + )?; + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.content_part.added", "item_id": message_id, "output_index": output_index, "content_index": 0, "part": {"type": "output_text", "text": "", "annotations": []}}), + )?; + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.output_text.delta", "item_id": message_id, "output_index": output_index, "content_index": 0, "delta": content}), + )?; + } + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.output_text.done", "item_id": message_id, "output_index": output_index, "content_index": 0, "text": content}), + )?; + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.content_part.done", "item_id": message_id, "output_index": output_index, "content_index": 0, "part": {"type": "output_text", "text": content, "annotations": []}}), + )?; + let item = json!({"id": message_id, "type": "message", "status": status, "role": "assistant", "content": [{"type": "output_text", "text": content, "annotations": []}]}); + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.output_item.done", "output_index": output_index, "item": item}), + )?; + terminal_items.push(item); + output_index += 1; + } + for call in &calls { + let item_id = random_id("fc_"); + let mut item = json!({"id": item_id, "type": "function_call", "status": status, "name": call.function.name, "call_id": call.id, "arguments": call.function.arguments}); + let mut added = item.clone(); + added["status"] = Value::String("in_progress".into()); + added["arguments"] = Value::String(String::new()); + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.output_item.added", "output_index": output_index, "item": added}), + )?; + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.function_call_arguments.delta", "item_id": item_id, "output_index": output_index, "delta": call.function.arguments}), + )?; + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.function_call_arguments.done", "item_id": item_id, "output_index": output_index, "name": call.function.name, "arguments": call.function.arguments}), + )?; + item["id"] = Value::String(item_id); + send_responses_sse( + stream, + &mut sequence, + json!({"type": "response.output_item.done", "output_index": output_index, "item": item}), + )?; + terminal_items.push(item); + output_index += 1; + } + let event_type = if finish == "length" { + "response.incomplete" + } else if finish == "error" { + "response.failed" + } else { + "response.completed" + }; + let cached = output.cached_tokens.min(output.prompt_tokens); + return send_responses_sse( + stream, + &mut sequence, + json!({"type": event_type, "response": {"id": id, "object": "response", "created_at": created, "status": status, "model": request.model_id, "output": terminal_items, "usage": {"input_tokens": output.prompt_tokens, "input_tokens_details": {"cached_tokens": cached, "cache_write_tokens": output.prompt_tokens - cached}, "output_tokens": output.completion_tokens, "output_tokens_details": {"reasoning_tokens": 0}, "total_tokens": output.prompt_tokens + output.completion_tokens}}}), + ); + } + GenerationEvent::Finished(Err(error)) => { + let _ = send_sse_error(stream, &error); + return Ok(()); + } + _ => {} + } + } + Err((500, "The model runtime stopped unexpectedly.".into())) +} + +#[allow(clippy::too_many_arguments)] +fn send_named_sse( + stream: &mut impl Write, + event: &str, + value: &Value, +) -> Result<(), (u16, String)> { + let body = serde_json::to_vec(value).map_err(|error| (500, error.to_string()))?; + stream + .write_all(b"event: ") + .and_then(|()| stream.write_all(event.as_bytes())) + .and_then(|()| stream.write_all(b"\ndata: ")) + .and_then(|()| stream.write_all(&body)) + .and_then(|()| stream.write_all(b"\n\n")) + .map_err(|error| (500, error.to_string())) +} + +fn send_responses_sse( + stream: &mut impl Write, + sequence: &mut u32, + value: Value, +) -> Result<(), (u16, String)> { + let mut object = value + .as_object() + .cloned() + .ok_or_else(|| (500, "Responses event is not an object".to_owned()))?; + let event_type = object.shift_remove("type").unwrap_or(Value::Null); + let mut ordered = Map::new(); + ordered.insert("type".into(), event_type); + ordered.insert("sequence_number".into(), Value::from(*sequence)); + ordered.extend(object); + *sequence += 1; + send_sse(stream, &Value::Object(ordered)).map_err(|error| (500, error)) +} + +fn stream_response_with_keepalive( + stream: &mut TcpStream, + state: &State, + request: ResponseOptions, + active: crate::runtime::ActiveGeneration, + id: &str, + keepalive_interval: Duration, +) -> Result<(), (u16, String)> { + let mut projector = ToolProjector::new(); + let mut output = None; + let mut prefilling = true; + let mut headers_sent = false; + let mut role_sent = false; + let mut last_keepalive = Instant::now(); + loop { + let event = match receive_stream_event( + stream, + &active, + prefilling && headers_sent, + &mut last_keepalive, + keepalive_interval, + ) { + Ok(Some(event)) => event, + Ok(None) => break, + Err(_) => return Ok(()), + }; + match event { + GenerationEvent::Loading => {} + GenerationEvent::Context { + tokens_per_second, .. + } => { + prefilling = tokens_per_second.is_none(); + if prefilling && !headers_sent { + send_sse_headers(stream).map_err(|error| (500, error))?; + headers_sent = true; + last_keepalive = Instant::now(); + } else if !prefilling { + send_stream_start(stream, &request, id, &mut headers_sent, &mut role_sent)?; + } + } + GenerationEvent::Chunk { reasoning, content } => { + prefilling = false; + send_stream_start(stream, &request, id, &mut headers_sent, &mut role_sent)?; + if request.has_tools && !reasoning { + let events = projector.push(&content, false, "call_"); + if send_chat_projection_events(stream, &request, id, events).is_err() { + active.cancel.store(true, Ordering::Relaxed); + return Ok(()); + } + if TOOL_SYNTAXES + .iter() + .any(|syntax| projector.raw.contains(syntax.tool_end)) + { + active.cancel.store(true, Ordering::Relaxed); + } + } else if !content.is_empty() { + let field = if reasoning { + "reasoning_content" + } else { + "content" + }; + let chunk = chunk_json(id, &request.model_id, json!({field: content}), None); + if send_sse(stream, &chunk).is_err() { + active.cancel.store(true, Ordering::Relaxed); + return Ok(()); + } + } + } + GenerationEvent::Finished(result) => { + match result { + Ok(result) => { + send_stream_start(stream, &request, id, &mut headers_sent, &mut role_sent)?; + output = Some(result); + } + Err(error) => { + if headers_sent { + let _ = send_sse_error(stream, &error); + return Ok(()); + } + return Err((500, error)); + } + } + break; + } + } + } + let output = match output { + Some(output) => output, + None if headers_sent => { + let _ = send_sse_error(stream, "The model runtime stopped unexpectedly."); + return Ok(()); + } + None => return Err((500, "The model runtime stopped unexpectedly.".into())), + }; + if request.has_tools { + let events = projector.push("", true, "call_"); + send_chat_projection_events(stream, &request, id, events)?; + } + let (_, calls) = parse_generated_tools_with_ids( + state, + &output.message.content, + Protocol::Chat, + &projector.ids, + ); + if request.has_tools && projector.ids.is_empty() && !calls.is_empty() { + send_sse( + stream, + &chunk_json( + id, + &request.model_id, + json!({"tool_calls": tool_calls_json(&calls)}), + None, + ), + ) + .map_err(|error| (500, error))?; + } + let finish = if calls.is_empty() { + output.finish_reason + } else { + "tool_calls" + }; + send_sse( + stream, + &chunk_json(id, &request.model_id, json!({}), Some(finish)), + ) + .map_err(|error| (500, error))?; + if request.include_usage { + let usage = json!({ + "id": id, + "object": "chat.completion.chunk", + "created": unix_time(), + "model": request.model_id, + "choices": [], + "usage": usage_json(output.prompt_tokens, output.cached_tokens, output.completion_tokens), + }); + send_sse(stream, &usage).map_err(|error| (500, error))?; + } + stream + .write_all(b"data: [DONE]\n\n") + .map_err(|error| (500, error.to_string())) +} + +fn send_chat_projection_events( + stream: &mut impl Write, + request: &ResponseOptions, + response_id: &str, + events: Vec, +) -> Result<(), (u16, String)> { + for event in events { + let delta = match event { + ToolProjectionEvent::Text(content) if !content.is_empty() => { + json!({"content": content}) + } + ToolProjectionEvent::Start { index, id, name } => json!({"tool_calls": [{ + "index": index, "id": id, "type": "function", + "function": {"name": name, "arguments": ""} + }]}), + ToolProjectionEvent::Arguments { index, fragment } => json!({ + "tool_calls": [{"index": index, "function": {"arguments": fragment}}] + }), + ToolProjectionEvent::End { .. } | ToolProjectionEvent::Text(_) => continue, + }; + send_sse( + stream, + &chunk_json(response_id, &request.model_id, delta, None), + ) + .map_err(|error| (500, error))?; + } + Ok(()) +} + +pub(super) fn send_stream_start( + stream: &mut impl Write, + request: &ResponseOptions, + id: &str, + headers_sent: &mut bool, + role_sent: &mut bool, +) -> Result<(), (u16, String)> { + if !*headers_sent { + send_sse_headers(stream).map_err(|error| (500, error))?; + *headers_sent = true; + } + if !*role_sent { + let role = chunk_json(id, &request.model_id, json!({"role": "assistant"}), None); + send_sse(stream, &role).map_err(|error| (500, error))?; + *role_sent = true; + } + Ok(()) +} + +pub(super) fn receive_stream_event( + stream: &mut impl Write, + active: &crate::runtime::ActiveGeneration, + prefilling: bool, + last_keepalive: &mut Instant, + keepalive_interval: Duration, +) -> Result, String> { + if !prefilling { + return Ok(active.events.recv().ok()); + } + loop { + if last_keepalive.elapsed() >= keepalive_interval { + if let Err(error) = stream.write_all(b": prefill\n\n") { + active.cancel.store(true, Ordering::Relaxed); + return Err(error.to_string()); + } + *last_keepalive = Instant::now(); + } + let remaining = keepalive_interval.saturating_sub(last_keepalive.elapsed()); + match active.events.recv_timeout(remaining) { + Ok(event) => return Ok(Some(event)), + Err(std::sync::mpsc::RecvTimeoutError::Timeout) => {} + Err(std::sync::mpsc::RecvTimeoutError::Disconnected) => return Ok(None), + } + } +} + +fn wait_for_output( + active: crate::runtime::ActiveGeneration, +) -> Result { + let mut content = String::new(); + while let Ok(event) = active.events.recv() { + match event { + GenerationEvent::Chunk { + reasoning: false, + content: chunk, + } => { + content.push_str(&chunk); + if TOOL_SYNTAXES + .iter() + .any(|syntax| content.contains(syntax.tool_end)) + { + active.cancel.store(true, Ordering::Relaxed); + } + } + GenerationEvent::Finished(result) => { + return result.map_err(|error| (500, error)); + } + _ => {} + } + } + Err((500, "The model runtime stopped unexpectedly.".into())) +} + +pub(super) fn usage_json(prompt: u32, cached: u32, completion: u32) -> Value { + let cached = cached.min(prompt); + json!({ + "prompt_tokens": prompt, + "completion_tokens": completion, + "total_tokens": prompt + completion, + "prompt_tokens_details": { + "cached_tokens": cached, + "cache_write_tokens": prompt - cached + } + }) +} + +pub(super) fn chunk_json(id: &str, model: &str, delta: Value, finish: Option<&str>) -> Value { + json!({ + "id": id, + "object": "chat.completion.chunk", + "created": unix_time(), + "model": model, + "choices": [{"index": 0, "delta": delta, "finish_reason": finish}] + }) +} diff --git a/src/server/tools.rs b/src/server/tools.rs new file mode 100644 index 0000000..c648e2f --- /dev/null +++ b/src/server/tools.rs @@ -0,0 +1,733 @@ +use super::*; + +enum ToolProjectionState { + Seeking, + Invokes, + Parameters, + Value, + Done, + Failed, +} + +pub(super) enum ToolProjectionEvent { + Text(String), + Start { + index: usize, + id: String, + name: String, + }, + Arguments { + index: usize, + fragment: String, + }, + End { + index: usize, + }, +} + +pub(super) struct ToolProjector { + pub(super) raw: String, + position: usize, + text_emitted: usize, + state: ToolProjectionState, + index: usize, + pub(super) ids: Vec, + first_parameter: bool, + string_parameter: bool, + syntax: Option, +} + +impl ToolProjector { + pub(super) fn new() -> Self { + Self { + raw: String::new(), + position: 0, + text_emitted: 0, + state: ToolProjectionState::Seeking, + index: 0, + ids: Vec::new(), + first_parameter: true, + string_parameter: false, + syntax: None, + } + } + + pub(super) fn push( + &mut self, + chunk: &str, + final_chunk: bool, + prefix: &str, + ) -> Vec { + self.raw.push_str(chunk); + let mut events = Vec::new(); + loop { + match self.state { + ToolProjectionState::Seeking => { + if let Some((start, syntax)) = TOOL_SYNTAXES + .iter() + .filter_map(|syntax| { + self.raw + .find(syntax.tool_start) + .map(|start| (start, *syntax)) + }) + .min_by_key(|(start, _)| *start) + { + if start > self.text_emitted { + let text = self.raw[self.text_emitted..start].trim_end(); + if !text.is_empty() { + events.push(ToolProjectionEvent::Text(text.to_owned())); + } + } + self.position = start + syntax.tool_start.len(); + self.text_emitted = start; + self.syntax = Some(syntax); + self.state = ToolProjectionState::Invokes; + } else { + let limit = if final_chunk { + self.raw.len() + } else { + TOOL_SYNTAXES + .iter() + .map(|syntax| { + safe_before_partial_marker(&self.raw, syntax.tool_start) + }) + .min() + .unwrap_or(self.raw.len()) + }; + if limit > self.text_emitted { + let text = &self.raw[self.text_emitted..limit]; + if !text.trim().is_empty() { + events.push(ToolProjectionEvent::Text(text.to_owned())); + self.text_emitted = limit; + } + } + break; + } + } + ToolProjectionState::Invokes => { + let syntax = self.syntax.unwrap(); + self.skip_whitespace(); + if self.full_at(syntax.tool_end) { + self.position += syntax.tool_end.len(); + self.state = ToolProjectionState::Done; + break; + } + if self.partial_at(syntax.tool_end) || self.partial_at(syntax.invoke_start) { + break; + } + if !self.full_at(syntax.invoke_start) { + self.state = ToolProjectionState::Failed; + break; + } + let Some(tag_end) = self.raw[self.position..].find('>') else { + break; + }; + let tag_end = self.position + tag_end + 1; + let Some(name) = dsml_attribute(&self.raw[self.position..tag_end], "name") + else { + self.state = ToolProjectionState::Failed; + break; + }; + let id = random_tool_id(prefix); + self.ids.push(id.clone()); + events.push(ToolProjectionEvent::Start { + index: self.index, + id, + name, + }); + events.push(ToolProjectionEvent::Arguments { + index: self.index, + fragment: "{".into(), + }); + self.position = tag_end; + self.first_parameter = true; + self.state = ToolProjectionState::Parameters; + } + ToolProjectionState::Parameters => { + let syntax = self.syntax.unwrap(); + self.skip_whitespace(); + if self.full_at(syntax.invoke_end) { + events.push(ToolProjectionEvent::Arguments { + index: self.index, + fragment: "}".into(), + }); + events.push(ToolProjectionEvent::End { index: self.index }); + self.position += syntax.invoke_end.len(); + self.index += 1; + self.state = ToolProjectionState::Invokes; + continue; + } + if self.partial_at(syntax.invoke_end) || self.partial_at(syntax.parameter_start) + { + break; + } + if !self.full_at(syntax.parameter_start) { + self.state = ToolProjectionState::Failed; + break; + } + let Some(tag_end) = self.raw[self.position..].find('>') else { + break; + }; + let tag_end = self.position + tag_end + 1; + let tag = &self.raw[self.position..tag_end]; + let Some(name) = dsml_attribute(tag, "name") else { + self.state = ToolProjectionState::Failed; + break; + }; + self.string_parameter = + dsml_attribute(tag, "string").as_deref() != Some("false"); + let mut fragment = if self.first_parameter { + String::new() + } else { + ",".into() + }; + self.first_parameter = false; + fragment + .push_str(&serde_json::to_string(&name).unwrap_or_else(|_| "\"\"".into())); + fragment.push(':'); + if self.string_parameter { + fragment.push('"'); + } + events.push(ToolProjectionEvent::Arguments { + index: self.index, + fragment, + }); + self.position = tag_end; + self.state = ToolProjectionState::Value; + } + ToolProjectionState::Value => { + let syntax = self.syntax.unwrap(); + if let Some(relative_end) = self.raw[self.position..].find(syntax.parameter_end) + { + let end = self.position + relative_end; + self.emit_value(end, &mut events); + if self.string_parameter { + events.push(ToolProjectionEvent::Arguments { + index: self.index, + fragment: "\"".into(), + }); + } + self.position = end + syntax.parameter_end.len(); + self.state = ToolProjectionState::Parameters; + continue; + } + let limit = safe_parameter_value_limit( + &self.raw, + self.position, + syntax.parameter_end, + self.string_parameter, + ); + self.emit_value(limit, &mut events); + break; + } + ToolProjectionState::Done | ToolProjectionState::Failed => break, + } + } + events + } + + fn emit_value(&mut self, end: usize, events: &mut Vec) { + if end <= self.position { + return; + } + let raw = &self.raw[self.position..end]; + let fragment = if self.string_parameter { + let value = unescape_dsml(raw); + let encoded = serde_json::to_string(&value).unwrap_or_else(|_| "\"\"".into()); + encoded[1..encoded.len() - 1].to_owned() + } else { + raw.to_owned() + }; + events.push(ToolProjectionEvent::Arguments { + index: self.index, + fragment, + }); + self.position = end; + } + + fn skip_whitespace(&mut self) { + while self.raw[self.position..] + .chars() + .next() + .is_some_and(char::is_whitespace) + { + self.position += self.raw[self.position..].chars().next().unwrap().len_utf8(); + } + } + + fn full_at(&self, marker: &str) -> bool { + self.raw.as_bytes()[self.position..].starts_with(marker.as_bytes()) + } + + fn partial_at(&self, marker: &str) -> bool { + let tail = &self.raw.as_bytes()[self.position..]; + tail.len() < marker.len() && marker.as_bytes().starts_with(tail) + } +} + +fn dsml_attribute(tag: &str, name: &str) -> Option { + let start = tag.find(&format!("{name}=\""))? + name.len() + 2; + let end = start + tag[start..].find('"')?; + Some(unescape_dsml(&tag[start..end])) +} + +fn safe_before_partial_marker(text: &str, marker: &str) -> usize { + let mut limit = text.len().saturating_sub(marker.len().saturating_sub(1)); + while !text.is_char_boundary(limit) { + limit -= 1; + } + limit +} + +fn safe_parameter_value_limit(text: &str, start: usize, end_marker: &str, string: bool) -> usize { + let bytes = text.as_bytes(); + let marker = end_marker.as_bytes(); + let mut limit = bytes.len(); + for length in (1..marker.len().min(bytes.len().saturating_sub(start) + 1)).rev() { + if bytes[start..].ends_with(&marker[..length]) { + limit -= length; + break; + } + } + if string { + for entity in ["&", "<", ">", """, "'"] { + let entity = entity.as_bytes(); + for length in 1..entity.len() { + if bytes[start..limit].ends_with(&entity[..length]) { + limit -= length; + return limit; + } + } + } + } + limit +} + +pub(super) fn render_messages( + state: &State, + messages: &[ApiMessage], + tools: &[Value], + tool_schemas: &[String], + tools_enabled: bool, + protocol: Protocol, +) -> Result<(String, Vec), (u16, String)> { + validate_tool_results(state, messages, protocol)?; + let preserve_reasoning = tools_enabled + || messages.iter().any(|message| { + matches!(message.role.as_str(), "tool" | "function") || !message.tool_calls.is_empty() + }); + let mut system = String::new(); + if tools_enabled { + system.push_str(TOOLS_PROMPT); + if tool_schemas.is_empty() { + for tool in tools { + let schema = tool.get("function").unwrap_or(tool); + if !system.ends_with("\n\n") { + system.push('\n'); + } + system.push_str( + &serde_json::to_string(schema) + .map_err(|error| (400, format!("invalid tool schema: {error}")))?, + ); + system.push('\n'); + } + } else { + for schema in tool_schemas { + if !system.ends_with("\n\n") { + system.push('\n'); + } + system.push_str(schema); + system.push('\n'); + } + } + system.push_str( + "\nYou MUST strictly follow the above defined tool name and parameter schemas to invoke tool calls. Use the exact parameter names from the schemas.", + ); + } + + let mut turns = Vec::::new(); + for message in messages { + let content = content_text(&message.content); + match message.role.as_str() { + "system" | "developer" => { + if !system.is_empty() { + system.push_str("\n\n"); + } + system.push_str(&content); + } + "user" => turns.push(ChatTurn { + user: true, + skip_previous_eos: false, + reasoning: None, + reasoning_complete: true, + content, + }), + "tool" | "function" => { + let wrapped = format!( + "{}", + escape_tool_result(&content) + ); + if let Some(previous) = turns.last_mut() + && previous.user + && previous.content.starts_with("") + { + previous.content.push_str(&wrapped); + } else { + turns.push(ChatTurn { + user: true, + skip_previous_eos: protocol == Protocol::Responses, + reasoning: None, + reasoning_complete: true, + content: wrapped, + }); + } + } + "assistant" => { + let mut content = content; + if !message.tool_calls.is_empty() { + content.push_str(&replayed_or_canonical_tools(state, &message.tool_calls)); + } + let reasoning = content_text(&message.reasoning_content); + turns.push(ChatTurn { + user: false, + skip_previous_eos: false, + reasoning: (preserve_reasoning && !reasoning.is_empty()).then_some(reasoning), + reasoning_complete: true, + content, + }); + } + role => return Err((400, format!("unsupported message role: {role}"))), + } + } + Ok((system, turns)) +} + +pub(super) fn validate_tool_results( + state: &State, + messages: &[ApiMessage], + protocol: Protocol, +) -> Result<(), (u16, String)> { + if !matches!(protocol, Protocol::Anthropic | Protocol::Responses) { + return Ok(()); + } + let memory = state.tool_memory.lock().ok(); + for (index, message) in messages.iter().enumerate() { + if !matches!(message.role.as_str(), "tool" | "function") || message.tool_call_id.is_empty() + { + continue; + } + let id = &message.tool_call_id; + let live = memory + .as_ref() + .is_some_and(|memory| memory.contains_key(id)); + let replayed = messages[..index].iter().any(|message| { + message.role == "assistant" && message.tool_calls.iter().any(|call| call.id == *id) + }); + if live || replayed { + continue; + } + let message = match protocol { + Protocol::Anthropic => format!( + "Anthropic continuation state is not available for tool_use_id {id}; retry by replaying the full messages history" + ), + Protocol::Responses => format!( + "Responses continuation state is not available for call_id {id}; retry by replaying the full input history" + ), + _ => unreachable!(), + }; + return Err((400, message)); + } + Ok(()) +} + +fn replayed_or_canonical_tools(state: &State, calls: &[ApiToolCall]) -> String { + if let Ok(memory) = state.tool_memory.lock() + && let Some(raw) = calls.iter().find_map(|call| { + (!call.id.is_empty()) + .then(|| memory.get(&call.id)) + .flatten() + }) + { + return raw.clone(); + } + canonical_tools(calls) +} + +pub(super) fn canonical_tools(calls: &[ApiToolCall]) -> String { + let mut output = String::from("\n\n<|DSML|tool_calls>\n"); + for call in calls { + output.push_str("<|DSML|invoke name=\""); + output.push_str(&escape_attribute(&call.function.name)); + output.push_str("\">\n"); + match serde_json::from_str::(&call.function.arguments) { + Ok(Value::Object(arguments)) => { + for (name, value) in arguments { + output.push_str("<|DSML|parameter name=\""); + output.push_str(&escape_attribute(&name)); + let string = value.as_str(); + output.push_str(if string.is_some() { + "\" string=\"true\">" + } else { + "\" string=\"false\">" + }); + if let Some(value) = string { + output.push_str(&escape_parameter(value)); + } else { + output.push_str(&escape_json_parameter(&value.to_string())); + } + output.push_str("\n"); + } + } + _ => { + output.push_str("<|DSML|parameter name=\"arguments\" string=\"true\">"); + output.push_str(&escape_parameter(&call.function.arguments)); + output.push_str("\n"); + } + } + output.push_str("\n"); + } + output.push_str(""); + output +} + +pub(super) fn parse_generated_tools( + state: &State, + text: &str, + protocol: Protocol, +) -> (String, Vec) { + parse_generated_tools_with_ids(state, text, protocol, &[]) +} + +pub(super) fn parse_generated_tools_with_ids( + state: &State, + text: &str, + protocol: Protocol, + streamed_ids: &[String], +) -> (String, Vec) { + let Some((start, syntax)) = TOOL_SYNTAXES + .iter() + .filter_map(|syntax| text.find(syntax.tool_start).map(|start| (start, *syntax))) + .min_by_key(|(start, _)| *start) + else { + return (text.to_owned(), Vec::new()); + }; + let Some(relative_end) = text[start..].find(syntax.tool_end) else { + return (text.to_owned(), Vec::new()); + }; + let end = start + relative_end + syntax.tool_end.len(); + let content = text[..start].trim_end(); + let raw = &text[content.len()..end]; + let mut calls = Vec::new(); + let mut cursor = raw.find(syntax.tool_start).unwrap() + syntax.tool_start.len(); + loop { + skip_text_whitespace(raw, &mut cursor); + if raw[cursor..].starts_with(syntax.tool_end) { + break; + } + if !raw[cursor..].starts_with(syntax.invoke_start) { + return (text.to_owned(), Vec::new()); + } + let Some(tag_end) = raw[cursor..].find('>').map(|end| cursor + end + 1) else { + return (text.to_owned(), Vec::new()); + }; + let Some(name) = dsml_attribute(&raw[cursor..tag_end], "name") else { + return (text.to_owned(), Vec::new()); + }; + cursor = tag_end; + let mut arguments = Map::new(); + loop { + skip_text_whitespace(raw, &mut cursor); + if raw[cursor..].starts_with(syntax.invoke_end) { + cursor += syntax.invoke_end.len(); + break; + } + let Some((name, value)) = parse_tool_parameter(raw, &mut cursor, syntax) else { + return (text.to_owned(), Vec::new()); + }; + arguments.insert(name, value); + } + calls.push(ApiToolCall { + id: String::new(), + function: ApiFunction { + name, + arguments: Value::Object(arguments).to_string(), + }, + }); + } + if calls.is_empty() { + return (text.to_owned(), Vec::new()); + } + let prefix = if protocol == Protocol::Anthropic { + "toolu_" + } else { + "call_" + }; + for (index, call) in calls.iter_mut().enumerate() { + call.id = streamed_ids + .get(index) + .cloned() + .unwrap_or_else(|| random_tool_id(prefix)); + } + if let Ok(mut memory) = state.tool_memory.lock() { + // ponytail: one process-local replay table; add LRU eviction if 100k live tool ids is measured insufficient. + if memory.len() >= 100_000 { + memory.clear(); + } + for call in &calls { + memory.insert(call.id.clone(), raw.to_owned()); + } + } + (content.to_owned(), calls) +} + +#[derive(Clone, Copy)] +pub(super) struct ToolSyntax { + pub(super) tool_start: &'static str, + pub(super) tool_end: &'static str, + pub(super) invoke_start: &'static str, + pub(super) invoke_end: &'static str, + pub(super) parameter_start: &'static str, + pub(super) parameter_end: &'static str, +} + +pub(super) const TOOL_SYNTAXES: [ToolSyntax; 3] = [ + ToolSyntax { + tool_start: "<|DSML|tool_calls>", + tool_end: "", + invoke_start: "<|DSML|invoke", + invoke_end: "", + parameter_start: "<|DSML|parameter", + parameter_end: "", + }, + ToolSyntax { + tool_start: "", + tool_end: "", + invoke_start: "", + parameter_start: "", + }, + ToolSyntax { + tool_start: "", + tool_end: "", + invoke_start: "", + parameter_start: "", + }, +]; + +fn parse_tool_parameter( + text: &str, + cursor: &mut usize, + syntax: ToolSyntax, +) -> Option<(String, Value)> { + if !text[*cursor..].starts_with(syntax.parameter_start) { + return None; + } + let tag_end = text[*cursor..].find('>').map(|end| *cursor + end + 1)?; + let tag = &text[*cursor..tag_end]; + let name = dsml_attribute(tag, "name")?; + let is_string = dsml_attribute(tag, "string"); + *cursor = tag_end; + let mut nested_start = *cursor; + skip_text_whitespace(text, &mut nested_start); + if is_string.is_none() && text[nested_start..].starts_with(syntax.parameter_start) { + *cursor = nested_start; + let mut nested = Map::new(); + loop { + skip_text_whitespace(text, cursor); + if !text[*cursor..].starts_with(syntax.parameter_start) { + break; + } + let (name, value) = parse_tool_parameter(text, cursor, syntax)?; + nested.insert(name, value); + } + skip_text_whitespace(text, cursor); + if !text[*cursor..].starts_with(syntax.parameter_end) { + return None; + } + *cursor += syntax.parameter_end.len(); + return Some((name, Value::Object(nested))); + } + let value_end = text[*cursor..] + .find(syntax.parameter_end) + .map(|end| *cursor + end)?; + let raw = &text[*cursor..value_end]; + *cursor = value_end + syntax.parameter_end.len(); + let value = if is_string.as_deref().unwrap_or("true") == "true" { + Value::String(unescape_dsml(raw)) + } else { + serde_json::from_str(raw).unwrap_or(Value::Null) + }; + Some((name, value)) +} + +fn skip_text_whitespace(text: &str, cursor: &mut usize) { + while text[*cursor..] + .chars() + .next() + .is_some_and(char::is_whitespace) + { + *cursor += text[*cursor..].chars().next().unwrap().len_utf8(); + } +} + +pub(super) fn tool_calls_json(calls: &[ApiToolCall]) -> Value { + Value::Array( + calls + .iter() + .map(|call| { + json!({ + "id": call.id, + "type": "function", + "function": { + "name": call.function.name, + "arguments": call.function.arguments, + } + }) + }) + .collect(), + ) +} + +pub(super) fn content_text(value: &Value) -> String { + match value { + Value::String(text) => text.clone(), + Value::Array(parts) => parts + .iter() + .filter_map(|part| match part { + Value::String(text) => Some(text.as_str()), + Value::Object(object) => object.get("text").and_then(Value::as_str), + _ => None, + }) + .collect(), + _ => String::new(), + } +} + +pub(super) fn escape_attribute(text: &str) -> String { + text.replace('&', "&") + .replace('<', "<") + .replace('>', ">") + .replace('"', """) +} + +pub(super) fn escape_parameter(text: &str) -> String { + text.replace("", "</|DSML|parameter>") +} + +pub(super) fn escape_json_parameter(text: &str) -> String { + text.replace("", "\\u003c/|DSML|parameter>") +} + +pub(super) fn escape_tool_result(text: &str) -> String { + text.replace("", "</tool_result>") +} + +pub(super) fn unescape_dsml(text: &str) -> String { + text.replace(""", "\"") + .replace(">", ">") + .replace("<", "<") + .replace("&", "&") +}