use crate::{
config::{self, Config, SharedConfig},
engine::Engine,
sys,
types::*,
update, AGENT_VERSION,
};
use anyhow::{anyhow, Result};
use chrono::{DateTime, SecondsFormat, Utc};
use parking_lot::Mutex;
use std::{
collections::VecDeque,
sync::{
atomic::{AtomicBool, Ordering},
Arc,
},
time::Duration,
};
use tracing::{info, warn};
const TRACE_TARGET: &str = "studio_worker::runtime";
pub const RECENT_JOBS_CAP: usize = 50;
pub const RECENT_LOGS_CAP: usize = 1000;
pub const PROMPT_PREVIEW_CHARS: usize = 200;
pub const LOG_SHIP_QUEUE_CAP: usize = 5_000;
#[derive(Debug, Clone, Copy, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
#[serde(rename_all = "snake_case")]
pub enum JobSource {
#[default]
Studio,
Local,
Lane,
Stream,
}
impl JobSource {
pub fn as_str(&self) -> &'static str {
match self {
JobSource::Studio => "studio",
JobSource::Local => "local",
JobSource::Lane => "lane",
JobSource::Stream => "stream",
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct CurrentJob {
pub job_id: String,
pub kind: TaskKind,
pub model: String,
pub prompt: String,
pub started_at: DateTime<Utc>,
pub source: JobSource,
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum JobOutcome {
Completed,
Failed { reason: String },
}
#[derive(Debug, Clone, PartialEq)]
pub struct RecentJob {
pub job_id: String,
pub kind: TaskKind,
pub model: String,
pub prompt: String,
pub outcome: JobOutcome,
pub started_at: DateTime<Utc>,
pub finished_at: DateTime<Utc>,
pub source: JobSource,
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
#[serde(tag = "outcome", rename_all = "snake_case")]
pub enum HeartbeatOutcome {
Ok,
Err { reason: String },
}
#[derive(Debug, Clone, PartialEq, Eq, Default, serde::Serialize, serde::Deserialize)]
#[serde(tag = "state", rename_all = "snake_case")]
pub enum SessionState {
#[default]
WaitingForApproval,
Connecting,
Connected,
Reconnecting { attempt: u32 },
AuthFailed { reason: String },
Fatal { reason: String },
Stopped,
}
impl SessionState {
pub fn summary(&self) -> String {
match self {
SessionState::WaitingForApproval => "waiting for studio operator approval".into(),
SessionState::Connecting => "connecting to the studio…".into(),
SessionState::Connected => "connected — ready for jobs".into(),
SessionState::Reconnecting { attempt } => {
format!("reconnecting (attempt {attempt})…")
}
SessionState::AuthFailed { reason } => format!(
"authentication failed: {reason}. Re-register with \
`studio-worker register --reset`."
),
SessionState::Fatal { reason } => {
format!("session ended: {reason}")
}
SessionState::Stopped => "stopped".into(),
}
}
}
#[derive(Debug, Clone, PartialEq, serde::Serialize, serde::Deserialize)]
#[serde(rename_all = "camelCase")]
pub struct HeartbeatStatus {
#[serde(flatten)]
pub outcome: HeartbeatOutcome,
pub last_attempt_at: DateTime<Utc>,
}
#[derive(Clone, Default)]
pub struct WorkerObservers {
pub current_job: Arc<Mutex<Option<CurrentJob>>>,
pub active_jobs: Arc<Mutex<Vec<CurrentJob>>>,
pub thumbnails: crate::thumbnail::Thumbnails,
pub recent_jobs: Arc<Mutex<VecDeque<RecentJob>>>,
pub local_jobs: Arc<Mutex<VecDeque<RecentJob>>>,
pub local_api_url: Arc<Mutex<Option<String>>>,
pub last_heartbeat: Arc<Mutex<Option<HeartbeatStatus>>>,
pub session_state: Arc<Mutex<SessionState>>,
pub gpu_runtime: Arc<Mutex<Option<GpuRuntimeStatus>>>,
pub catalog: Arc<Mutex<crate::catalog::Catalog>>,
pub catalog_path: Arc<Mutex<Option<std::path::PathBuf>>>,
pub recent_logs: Arc<Mutex<VecDeque<LogEntry>>>,
pub recent_logs_seq: Arc<std::sync::atomic::AtomicU64>,
}
pub fn recent_logs_after(observers: &WorkerObservers, after: u64) -> (Vec<LogEntry>, u64) {
let ring = observers.recent_logs.lock();
let newest = observers.recent_logs_seq.load(Ordering::SeqCst);
let oldest = newest - ring.len() as u64;
let after = if after > newest { 0 } else { after.max(oldest) };
let skip = (after - oldest) as usize;
(ring.iter().skip(skip).cloned().collect(), newest)
}
impl WorkerObservers {
pub fn with_global_worker_log() -> Self {
let ring = crate::job_log::global_worker_log();
Self {
recent_logs: ring.entries.clone(),
recent_logs_seq: ring.seq.clone(),
..Self::default()
}
}
pub fn worker_log(&self) -> crate::job_log::WorkerLogRing {
crate::job_log::WorkerLogRing {
entries: self.recent_logs.clone(),
seq: self.recent_logs_seq.clone(),
}
}
}
pub fn set_session_state(observers: &WorkerObservers, state: SessionState) {
*observers.session_state.lock() = state;
}
#[derive(Debug, Clone, PartialEq, Eq, serde::Serialize, serde::Deserialize)]
pub struct GpuRuntimeStatus {
pub ok: bool,
pub detail: String,
}
pub fn sync_studio_model(
observers: &WorkerObservers,
model_id: &str,
kind: TaskKind,
source: &ModelSource,
) {
let incoming = crate::catalog::CatalogModel {
id: model_id.to_string(),
display_name: model_id.to_string(),
kind,
vram_gb_estimate: 0.0,
description: None,
source: source.clone(),
enabled: true,
origin: "studio".into(),
exclusive_group: None,
};
let changed = observers.catalog.lock().sync_studio_model(incoming);
if !changed {
return;
}
let path = observers.catalog_path.lock().clone();
if let Some(path) = path {
let snapshot = observers.catalog.lock().clone();
if let Err(e) = snapshot.save(&path) {
warn!(
target: TRACE_TARGET,
op = "catalog_sync",
model_id,
error = %e,
"failed to persist studio model into the local catalog"
);
} else {
info!(
target: TRACE_TARGET,
op = "catalog_sync",
model_id,
kind = kind.as_str(),
"mirrored studio model into the local catalog"
);
}
}
}
pub fn set_gpu_runtime_status(observers: &WorkerObservers, status: Result<()>) {
let value = match &status {
Ok(()) => GpuRuntimeStatus {
ok: true,
detail: "GPU runtime available".into(),
},
Err(e) => {
warn!(
target: TRACE_TARGET,
op = "gpu_preflight",
error = %e,
"GPU runtime missing at startup; image jobs will fail until it is installed"
);
GpuRuntimeStatus {
ok: false,
detail: e.to_string(),
}
}
};
*observers.gpu_runtime.lock() = Some(value);
}
pub fn truncate_prompt(s: &str) -> String {
if s.chars().count() <= PROMPT_PREVIEW_CHARS {
return s.to_string();
}
let mut out: String = s.chars().take(PROMPT_PREVIEW_CHARS).collect();
out.push('…');
out
}
pub fn record_recent_job(observers: &WorkerObservers, entry: RecentJob) {
let mut ring = observers.recent_jobs.lock();
ring.push_front(entry);
while ring.len() > RECENT_JOBS_CAP {
ring.pop_back();
}
}
pub fn record_local_job(observers: &WorkerObservers, entry: RecentJob) {
let mut ring = observers.local_jobs.lock();
ring.push_front(entry);
while ring.len() > RECENT_JOBS_CAP {
ring.pop_back();
}
}
#[doc(hidden)]
pub fn push_recent_job_for_tests(observers: &WorkerObservers, job_id: &str) {
let now = Utc::now();
record_recent_job(
observers,
RecentJob {
job_id: job_id.to_string(),
kind: TaskKind::Image,
model: "synthetic".into(),
prompt: String::new(),
outcome: JobOutcome::Completed,
started_at: now,
finished_at: now,
source: JobSource::Studio,
},
);
}
pub const AUTO_UPDATE_TICK: Duration = Duration::from_secs(60);
pub const AUTO_UPDATE_SHUTDOWN_TICK: Duration = Duration::from_millis(250);
pub const HEARTBEAT_INTERVAL: Duration = Duration::from_secs(5);
#[derive(Debug, Clone, Copy)]
pub struct LoopSchedule {
pub ws_session: crate::ws::session::SessionSchedule,
pub auto_update_tick: Duration,
pub shutdown_tick: Duration,
}
impl Default for LoopSchedule {
fn default() -> Self {
Self {
ws_session: crate::ws::session::SessionSchedule::default(),
auto_update_tick: AUTO_UPDATE_TICK,
shutdown_tick: AUTO_UPDATE_SHUTDOWN_TICK,
}
}
}
impl LoopSchedule {
pub fn fast_for_tests() -> Self {
Self {
ws_session: crate::ws::session::SessionSchedule::fast_for_tests(),
auto_update_tick: Duration::from_millis(1),
shutdown_tick: Duration::from_millis(1),
}
}
}
#[derive(Debug, Clone, Default)]
pub struct RegisterArgs {
pub api_base_url: Option<String>,
pub reset: bool,
}
pub async fn register(config_path: Option<&str>, args: RegisterArgs) -> Result<()> {
let (mut cfg, path) = config::load(config_path)?;
if args.reset {
clear_registration(&mut cfg);
}
if let Some(url) = args.api_base_url {
cfg.api_base_url = url;
}
config::save(&cfg, &path)?;
if args.reset {
info!(
config_path = %path.display(),
"local registration state cleared; next launch will auto-register"
);
println!(
"local registration state cleared; run `studio-worker run` or \
`studio-worker ui` to auto-register"
);
} else {
info!(
config_path = %path.display(),
"register flags persisted; next launch will auto-register"
);
println!(
"saved; run `studio-worker run` or `studio-worker ui` to auto-register against {}",
cfg.api_base_url
);
}
Ok(())
}
pub async fn status(config_path: Option<&str>) -> Result<()> {
let (cfg, path) = config::load(config_path)?;
println!("{}", format_status(&cfg, &path));
Ok(())
}
pub fn format_status(cfg: &Config, path: &std::path::Path) -> String {
let mut out = String::new();
use std::fmt::Write as _;
let _ = writeln!(out, "config path: {}", path.display());
let _ = writeln!(out, "api_base_url: {}", cfg.api_base_url);
let registration_line = if cfg.worker_id.is_some() && cfg.auth_token.is_some() {
format!("approved as {}", cfg.worker_id.as_deref().unwrap_or(""))
} else if let Some(rid) = cfg.registration_request_id.as_deref() {
format!("pending operator approval (request {rid})")
} else {
"not registered (will auto-register on next launch)".into()
};
let _ = writeln!(out, "registration: {registration_line}");
let _ = writeln!(out, "vram_threshold_gb: {}", cfg.vram_threshold_gb);
let _ = writeln!(out, "models_root: {}", cfg.models_root.display());
let _ = writeln!(out, "auto_update: {}", cfg.auto_update_enabled);
let _ = writeln!(
out,
"update_interval: {}s",
cfg.auto_update_interval_secs
);
out
}
pub fn set_threshold(config_path: Option<&str>, gb: f32) -> Result<()> {
if gb < 0.0 {
return Err(anyhow!("threshold must be >= 0"));
}
let (mut cfg, path) = config::load(config_path)?;
cfg.vram_threshold_gb = gb;
config::save(&cfg, &path)?;
info!(
target: TRACE_TARGET,
op = "set_threshold",
vram_threshold_gb = gb,
config_path = path.display().to_string(),
"VRAM threshold persisted"
);
println!("vram_threshold_gb = {gb}");
Ok(())
}
pub fn log_startup_banner(cfg: &Config, path: &std::path::Path) {
info!(
target: TRACE_TARGET,
op = "startup",
version = AGENT_VERSION,
config_path = path.display().to_string(),
api_base_url = cfg.api_base_url.as_str(),
vram_threshold_gb = cfg.vram_threshold_gb,
auto_update_enabled = cfg.auto_update_enabled,
auto_update_interval_secs = cfg.auto_update_interval_secs,
models_root = cfg.models_root.display().to_string(),
worker_id = cfg.worker_id.as_deref().unwrap_or("(unregistered)"),
"studio-worker booting"
);
}
pub fn show_config(config_path: Option<&str>) -> Result<()> {
let (cfg, path) = config::load(config_path)?;
println!("# {}", path.display());
print!("{}", toml::to_string_pretty(&cfg)?);
Ok(())
}
pub async fn check_update(config_path: Option<&str>) -> Result<()> {
let (cfg, _) = config::load(config_path)?;
let current = semver::Version::parse(AGENT_VERSION)
.map_err(|e| anyhow!("invalid current version {AGENT_VERSION}: {e}"))?;
let outcome = tokio::task::spawn_blocking(move || {
update::check(&cfg.auto_update_feed, ¤t, cfg.auto_update_prerelease)
})
.await??;
println!("{}", format_check_outcome(&outcome));
Ok(())
}
pub fn format_check_outcome(outcome: &update::CheckOutcome) -> String {
match outcome {
update::CheckOutcome::UpToDate { current } => format!("up to date: {current}"),
update::CheckOutcome::NewerAvailable { current, latest } => {
format!("update available: {current} -> {latest}")
}
}
}
pub async fn run(config_path: Option<&str>, wait_for_lock: bool) -> Result<()> {
let (cfg, path) = config::load(config_path)?;
let _lock = match crate::daemon_lock::acquire(&path)? {
crate::daemon_lock::Acquired::Mine(lock) => lock,
crate::daemon_lock::Acquired::HeldElsewhere if wait_for_lock => {
let path = path.clone();
tokio::task::spawn_blocking(move || {
crate::daemon_lock::wait_until_acquired(&path, crate::daemon_lock::WAIT_POLL)
})
.await??
}
crate::daemon_lock::Acquired::HeldElsewhere => return Ok(()),
};
log_startup_banner(&cfg, &path);
let control = crate::control::DaemonControl::new(
config::shared(cfg),
path,
sys::detect_vram_gb().unwrap_or(0.0),
);
let busy = Arc::new(AtomicBool::new(false));
let logs: Arc<Mutex<Vec<LogEntry>>> = Arc::new(Mutex::new(Vec::new()));
let observers = WorkerObservers::with_global_worker_log();
let stop_clone = control.stop.clone();
tokio::spawn(async move {
let signal = wait_for_shutdown_signal().await;
request_shutdown(&stop_clone, signal);
});
let gate = crate::job_gate::JobGate::from_shared(busy.clone());
let local_api = spawn_local_api(&control, observers.clone(), gate);
let outcome = serve_studio(&control, logs, busy, observers, LoopSchedule::default()).await;
control.stop.store(true, Ordering::SeqCst);
if let Some(handle) = local_api {
let _ = handle.join();
}
outcome
}
pub const REGISTRATION_RESET_POLL: Duration = Duration::from_millis(250);
pub async fn serve_studio(
control: &crate::control::DaemonControl,
logs: Arc<Mutex<Vec<LogEntry>>>,
busy: Arc<AtomicBool>,
observers: WorkerObservers,
schedule: LoopSchedule,
) -> Result<()> {
loop {
match ensure_registered(
&control.cfg,
&control.config_path,
&control.registration,
&control.stop,
)
.await
{
Ok(RegistrationGate::Stopped) => {
info!(
target: TRACE_TARGET,
op = "shutdown",
"stopped before registration completed; exiting cleanly"
);
return Ok(());
}
Ok(RegistrationGate::Ready) => {
return run_loops(
control.cfg.clone(),
control.stop.clone(),
logs,
busy,
control.paused.clone(),
observers,
schedule,
)
.await;
}
Err(err) => {
tracing::error!(
target: TRACE_TARGET,
op = "registration",
error = %err,
"studio registration rejected; the local API keeps serving; \
reset the registration from the tray UI to ask again"
);
if !wait_for_registration_reset(control).await {
return Ok(());
}
reset_registration(control)?;
}
}
}
}
async fn wait_for_registration_reset(control: &crate::control::DaemonControl) -> bool {
loop {
if control.stop.load(Ordering::SeqCst) {
return false;
}
if control.reset_requested.swap(false, Ordering::SeqCst) {
return true;
}
tokio::time::sleep(REGISTRATION_RESET_POLL).await;
}
}
pub fn clear_registration(cfg: &mut Config) {
cfg.worker_id = None;
cfg.auth_token = None;
cfg.registration_request_id = None;
cfg.registration_secret = None;
cfg.install_id = None;
}
fn reset_registration(control: &crate::control::DaemonControl) -> Result<()> {
let snapshot = {
let mut cfg = control.cfg.lock();
clear_registration(&mut cfg);
cfg.clone()
};
config::save(&snapshot, &control.config_path)?;
*control.registration.lock() = crate::auto_register::RegistrationState::Pristine;
info!(
target: TRACE_TARGET,
op = "registration",
"registration reset; asking the studio again"
);
Ok(())
}
pub fn request_shutdown(stop: &AtomicBool, signal: &str) {
let already_stopping = stop.swap(true, Ordering::SeqCst);
info!(
target: TRACE_TARGET,
op = "shutdown",
signal,
already_stopping,
"shutdown signal received; stopping worker gracefully"
);
}
#[cfg_attr(coverage_nightly, coverage(off))]
async fn wait_for_shutdown_signal() -> &'static str {
#[cfg(unix)]
{
use tokio::signal::unix::{signal, SignalKind};
let mut sigterm = match signal(SignalKind::terminate()) {
Ok(s) => s,
Err(e) => {
warn!(
target: TRACE_TARGET,
op = "shutdown",
error = %e,
"could not install SIGTERM handler; falling back to Ctrl-C only"
);
let _ = tokio::signal::ctrl_c().await;
return "SIGINT";
}
};
tokio::select! {
_ = tokio::signal::ctrl_c() => "SIGINT",
_ = sigterm.recv() => "SIGTERM",
}
}
#[cfg(not(unix))]
{
let _ = tokio::signal::ctrl_c().await;
"ctrl-c"
}
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum RegistrationGate {
Ready,
Stopped,
}
pub async fn ensure_registered(
cfg: &SharedConfig,
path: &std::path::Path,
registration: &crate::auto_register::SharedRegistration,
stop: &Arc<AtomicBool>,
) -> Result<RegistrationGate> {
use std::time::Duration;
loop {
if stop.load(Ordering::SeqCst) {
return Ok(RegistrationGate::Stopped);
}
{
let snap = cfg.lock();
if snap.worker_id.is_some() && snap.auth_token.is_some() {
return Ok(RegistrationGate::Ready);
}
}
let state = crate::auto_register::tick(cfg, path, registration).await;
match state {
crate::auto_register::RegistrationState::Approved => {
return Ok(RegistrationGate::Ready)
}
crate::auto_register::RegistrationState::Rejected { reason } => {
return Err(anyhow!(
"registration rejected by the studio operator: {reason}. \
Run `studio-worker register --reset` to clear local state \
and submit a fresh request."
));
}
_ => {}
}
for _ in 0..30 {
if stop.load(Ordering::SeqCst) {
return Ok(RegistrationGate::Stopped);
}
tokio::time::sleep(Duration::from_secs(1)).await;
}
}
}
pub async fn run_loops(
cfg: SharedConfig,
stop: Arc<AtomicBool>,
logs: Arc<Mutex<Vec<LogEntry>>>,
busy: Arc<AtomicBool>,
paused: Arc<AtomicBool>,
observers: WorkerObservers,
schedule: LoopSchedule,
) -> Result<()> {
let session = crate::ws::session::spawn_ws_session(
cfg.clone(),
stop.clone(),
logs.clone(),
busy.clone(),
paused.clone(),
observers.clone(),
schedule.ws_session,
);
let auto_updater = spawn_auto_updater(
cfg.clone(),
stop.clone(),
logs.clone(),
busy.clone(),
schedule,
);
let (session_result, _) = tokio::join!(session, auto_updater);
session_result
}
pub const DEFAULT_LOCAL_API_PORT: u16 = 4787;
pub fn resolve_local_api_port(env_value: Option<&str>, cfg_port: Option<u16>) -> u16 {
resolve_port(
"STUDIO_WORKER_LOCAL_API_PORT",
env_value,
cfg_port,
DEFAULT_LOCAL_API_PORT,
)
}
pub fn resolve_stream_port(env_value: Option<&str>, cfg_port: Option<u16>) -> u16 {
resolve_port(
"STUDIO_WORKER_STREAM_PORT",
env_value,
cfg_port,
crate::stt_stream::server::DEFAULT_STREAM_PORT,
)
}
fn resolve_port(
env_name: &str,
env_value: Option<&str>,
cfg_port: Option<u16>,
default: u16,
) -> u16 {
let fallback = cfg_port.unwrap_or(default);
match env_value {
Some(raw) => match raw.parse::<u16>() {
Ok(port) => port,
Err(_) => {
warn!(
target: "studio_worker::local_api",
op = "resolve_port",
invalid = raw,
fallback,
"{env_name} is not a valid port; falling back"
);
fallback
}
},
None => fallback,
}
}
pub fn ensure_local_api_token(cfg: &SharedConfig, config_path: &std::path::Path) -> String {
let mut snap = cfg.lock();
if let Some(token) = snap.local_api_token.clone() {
return token;
}
let token = crate::secrets::new_secret_hex();
snap.local_api_token = Some(token.clone());
let snapshot = snap.clone();
drop(snap);
if let Err(e) = config::save(&snapshot, config_path) {
warn!(
target: "studio_worker::local_api",
op = "ensure_token",
config_path = %config_path.display(),
error = %e,
"failed to persist the local api token; a fresh one will be minted next launch"
);
}
token
}
#[cfg_attr(coverage_nightly, coverage(off))]
fn spawn_stream_listener(
cfg: &SharedConfig,
services: &crate::local_api::ModelServices,
observers: &WorkerObservers,
stop: Arc<AtomicBool>,
) {
let port = resolve_stream_port(
std::env::var("STUDIO_WORKER_STREAM_PORT").ok().as_deref(),
cfg.lock().stream_port,
);
let addr = format!("0.0.0.0:{port}");
match crate::stt_stream::server::StreamServer::bind(
&addr,
services.host.clone(),
services.tokens.clone(),
observers.clone(),
) {
Ok(server) => {
let bound = server.local_addr().port();
services.stream_port.store(bound, Ordering::SeqCst);
tracing::info!(target: "studio_worker::stt_stream", addr = %server.local_addr(), "stream listener listening");
std::thread::spawn(move || server.serve(&stop));
}
Err(err) => tracing::warn!(
target: "studio_worker::stt_stream",
%addr,
error = %err,
"stream listener could not bind; streaming speech is unavailable"
),
}
}
pub fn spawn_local_api(
control: &crate::control::DaemonControl,
observers: WorkerObservers,
gate: crate::job_gate::JobGate,
) -> Option<std::thread::JoinHandle<()>> {
let cfg = control.cfg.clone();
let config_path = control.config_path.as_path();
let stop = control.stop.clone();
let engine: Arc<dyn crate::engine::Engine> = match crate::engine::build(&cfg.lock()) {
Ok(engine) => engine.into(),
Err(err) => {
tracing::warn!(target: "studio_worker::local_api", error = %err, "local api: engine build failed");
return None;
}
};
let (loaded, catalog_path) =
crate::catalog::Catalog::load_for_serving(crate::config::catalog_path_for(config_path));
*observers.catalog.lock() = loaded;
*observers.catalog_path.lock() = catalog_path.clone();
let catalog = observers.catalog.clone();
let token = ensure_local_api_token(&cfg, config_path);
let port = resolve_local_api_port(
std::env::var("STUDIO_WORKER_LOCAL_API_PORT")
.ok()
.as_deref(),
cfg.lock().local_api_port,
);
let models_root = Some(cfg.lock().models_root.clone());
let host = crate::host::ModelHost::new(
catalog.clone(),
Arc::new(crate::loaders::Loaders::new(cfg.lock().models_root.clone())),
Arc::new(crate::admission::SystemProbe),
crate::residency::Residency::load_for_serving(crate::config::residency_path_for(
config_path,
)),
);
let services = crate::local_api::ModelServices::new(host.clone());
spawn_stream_listener(&cfg, &services, &observers, stop.clone());
let api = crate::local_api::LocalApi::bind(
&format!("127.0.0.1:{port}"),
engine.clone(),
catalog.clone(),
catalog_path.clone(),
observers.clone(),
token.clone(),
gate.clone(),
models_root.clone(),
services.clone(),
)
.or_else(|_| {
crate::local_api::LocalApi::bind(
"127.0.0.1:0",
engine,
catalog,
catalog_path,
observers.clone(),
token.clone(),
gate.clone(),
models_root,
services.clone(),
)
});
let api = match api {
Ok(api) => api.with_control(control.clone()),
Err(err) => {
tracing::warn!(target: "studio_worker::local_api", error = %err, "local api: bind failed");
return None;
}
};
set_gpu_runtime_status(
&observers,
crate::engine::sd_provision::vulkan_runtime_status(),
);
host.restore_residents();
let url = api.url();
*observers.local_api_url.lock() = Some(url.clone());
tracing::info!(target: "studio_worker::local_api", url = %url, "local image API listening");
let discovery_path = crate::config::local_api_discovery_path_for(config_path);
if let Some(path) = &discovery_path {
if let Err(e) = crate::local_api::write_discovery_file(path, &url, &token) {
tracing::warn!(
target: "studio_worker::local_api",
error = %e,
path = %path.display(),
"failed to write the local api discovery file"
);
}
}
Some(std::thread::spawn(move || {
api.serve(&stop);
if let Some(path) = &discovery_path {
crate::local_api::remove_discovery_file(path);
}
}))
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum AutoUpdateDecision {
Disabled,
SkippedBusy,
UpToDate,
CheckError(String),
Updated,
UpdateError(String),
}
pub async fn auto_update_tick(
cfg: &Config,
gate: &crate::job_gate::JobGate,
logs: &Arc<Mutex<Vec<LogEntry>>>,
) -> AutoUpdateDecision {
if !cfg.auto_update_enabled {
return AutoUpdateDecision::Disabled;
}
if gate.is_busy() {
push_log(
logs,
"info",
"auto-update",
"skipping check: worker is busy on a job",
None,
);
return AutoUpdateDecision::SkippedBusy;
}
let feed = cfg.auto_update_feed.clone();
let prerelease = cfg.auto_update_prerelease;
let logs_for_task = logs.clone();
let gate = gate.clone();
let outcome = tokio::task::spawn_blocking(move || -> Result<AutoUpdateDecision> {
let current = semver::Version::parse(AGENT_VERSION)
.map_err(|e| anyhow!("invalid AGENT_VERSION {AGENT_VERSION}: {e}"))?;
match update::check(&feed, ¤t, prerelease) {
Ok(update::CheckOutcome::UpToDate { current }) => {
push_log(
&logs_for_task,
"info",
"auto-update",
&format!("up to date at {current}"),
None,
);
Ok(AutoUpdateDecision::UpToDate)
}
Ok(update::CheckOutcome::NewerAvailable { current, latest }) => {
let Some(_reservation) = gate.try_reserve() else {
push_log(
&logs_for_task,
"info",
"auto-update",
"update available but a job started; deferring install",
None,
);
return Ok(AutoUpdateDecision::SkippedBusy);
};
push_log(
&logs_for_task,
"info",
"auto-update",
&format!("update available {current} -> {latest}; applying"),
None,
);
match update::apply(&feed, &latest) {
Ok(()) => {
push_log(
&logs_for_task,
"info",
"auto-update",
"binary replaced; restart pending",
None,
);
Ok(AutoUpdateDecision::Updated)
}
Err(e) => {
push_log(
&logs_for_task,
"error",
"auto-update",
&format!("update failed: {e}"),
None,
);
Ok(AutoUpdateDecision::UpdateError(e.to_string()))
}
}
}
Err(e) => {
push_log(
&logs_for_task,
"warn",
"auto-update",
&format!("check failed: {e}"),
None,
);
Ok(AutoUpdateDecision::CheckError(e.to_string()))
}
}
})
.await;
match outcome {
Ok(Ok(decision)) => decision,
Ok(Err(e)) => AutoUpdateDecision::CheckError(e.to_string()),
Err(e) => AutoUpdateDecision::CheckError(e.to_string()),
}
}
pub(crate) async fn wait_with_stop(total: Duration, stop: &Arc<AtomicBool>, tick: Duration) {
let mut elapsed = Duration::ZERO;
while elapsed < total {
if stop.load(Ordering::SeqCst) {
return;
}
let next = tick.min(total - elapsed);
tokio::time::sleep(next).await;
elapsed += next;
}
}
pub fn spawn_auto_updater(
cfg: SharedConfig,
stop: Arc<AtomicBool>,
logs: Arc<Mutex<Vec<LogEntry>>>,
busy: Arc<AtomicBool>,
schedule: LoopSchedule,
) -> tokio::task::JoinHandle<()> {
tokio::spawn(async move {
let mut elapsed = Duration::from_secs(0);
while !stop.load(Ordering::SeqCst) {
wait_with_stop(schedule.auto_update_tick, &stop, schedule.shutdown_tick).await;
if stop.load(Ordering::SeqCst) {
break;
}
elapsed += schedule.auto_update_tick;
let snapshot = cfg.lock().clone();
if elapsed < Duration::from_secs(snapshot.auto_update_interval_secs) {
continue;
}
elapsed = Duration::from_secs(0);
let gate = crate::job_gate::JobGate::from_shared(busy.clone());
let decision = auto_update_tick(&snapshot, &gate, &logs).await;
if matches!(decision, AutoUpdateDecision::Updated) {
stop.store(true, Ordering::SeqCst);
update::restart_self();
}
}
})
}
pub fn prompt_for(task: &Task) -> String {
match task {
Task::Image(p) => p.prompt.clone(),
Task::Llm(p) => p
.messages
.last()
.map(|m| m.content.clone())
.unwrap_or_default(),
Task::AudioStt(p) => p.input_url.clone(),
Task::AudioTts(p) => p.text.clone(),
Task::Video(p) => p.prompt.clone(),
}
}
pub fn is_unsupported_kind(e: &anyhow::Error) -> bool {
e.chain().any(|cause| {
cause
.downcast_ref::<crate::engine::UnsupportedTask>()
.is_some()
}) || e.to_string().contains("cannot serve")
}
pub fn build_capabilities(cfg: &Config, engine: &dyn Engine) -> WorkerCapabilities {
build_capabilities_with(cfg, engine, true)
}
pub fn build_capabilities_with(
cfg: &Config,
engine: &dyn Engine,
auto_enabled: bool,
) -> WorkerCapabilities {
let vram = sys::detect_vram_gb().unwrap_or(0.0);
let caps = engine.capabilities();
let supported_models_per_kind = caps.supported_models_per_kind.clone();
let task_kinds = caps.kinds();
let supported_models = {
let mut all = caps.flat_models();
all.sort();
all.dedup();
all
};
WorkerCapabilities {
machine_name: sys::machine_name(),
username: sys::username(),
agent_version: AGENT_VERSION.to_string(),
engine: engine.name().to_string(),
vram_total_gb: vram,
vram_threshold_gb: cfg.vram_threshold_gb,
auto_enabled,
auto_start: cfg!(feature = "ui"),
supported_models,
task_kinds,
supported_models_per_kind,
}
}
pub fn summarize_capabilities(caps: &WorkerCapabilities) -> String {
let kinds = caps
.task_kinds
.iter()
.map(|k| k.as_str())
.collect::<Vec<_>>()
.join(", ");
format!(
"advertising engine={}, vram={:.1}/{:.1}GB threshold, auto_enabled={}, \
kinds=[{}], {} model(s)=[{}]",
caps.engine,
caps.vram_total_gb,
caps.vram_threshold_gb,
caps.auto_enabled,
kinds,
caps.supported_models.len(),
caps.supported_models.join(", "),
)
}
pub fn vram_threshold_warning(caps: &WorkerCapabilities) -> Option<String> {
if caps.vram_total_gb > 0.0 && caps.vram_threshold_gb > caps.vram_total_gb {
Some(format!(
"configured VRAM threshold {:.1}GB exceeds detected GPU VRAM {:.1}GB; \
the studio may offer jobs larger than this card can fit and they will \
OOM on load — lower vram_threshold_gb to at or below {:.1}GB",
caps.vram_threshold_gb, caps.vram_total_gb, caps.vram_total_gb
))
} else {
None
}
}
pub fn push_log(
logs: &Arc<Mutex<Vec<LogEntry>>>,
level: &str,
category: &str,
message: &str,
job_id: Option<String>,
) {
push_log_with_observers(logs, None, level, category, message, job_id);
}
pub fn push_log_with_observers(
logs: &Arc<Mutex<Vec<LogEntry>>>,
observers: Option<&WorkerObservers>,
level: &str,
category: &str,
message: &str,
job_id: Option<String>,
) {
let entry = LogEntry {
ts: Utc::now().to_rfc3339_opts(SecondsFormat::Millis, true),
level: level.to_string(),
category: category.to_string(),
message: message.to_string(),
job_id,
};
let job_id = entry.job_id.as_deref();
if level == "error" {
tracing::error!(target: "studio_worker", job_id, "[{category}] {message}");
} else if level == "warn" {
tracing::warn!(target: "studio_worker", job_id, "[{category}] {message}");
} else {
info!(target: "studio_worker", job_id, "[{category}] {message}");
}
{
let mut queue = logs.lock();
if queue.len() >= LOG_SHIP_QUEUE_CAP {
let overflow = queue.len() + 2 - LOG_SHIP_QUEUE_CAP;
queue.drain(0..overflow);
queue.push(LogEntry {
ts: Utc::now().to_rfc3339_opts(SecondsFormat::Millis, true),
level: "warn".to_string(),
category: "logs".to_string(),
message: format!(
"ship queue full ({LOG_SHIP_QUEUE_CAP} entries); dropped {overflow} oldest"
),
job_id: None,
});
}
queue.push(entry.clone());
}
if let Some(o) = observers {
o.worker_log().push(entry);
}
}
pub fn restore_unshipped(logs: &Arc<Mutex<Vec<LogEntry>>>, mut batch: Vec<LogEntry>) {
let mut queue = logs.lock();
batch.append(&mut queue);
*queue = batch;
if queue.len() > LOG_SHIP_QUEUE_CAP {
let overflow = queue.len() - LOG_SHIP_QUEUE_CAP;
queue.drain(0..overflow);
}
}
#[cfg(test)]
mod tests {
use super::*;
use crate::config::Config;
use crate::engine::SyntheticEngine;
fn push_messages(observers: &WorkerObservers, count: usize) {
let logs = Arc::new(Mutex::new(Vec::new()));
for i in 0..count {
push_log_with_observers(&logs, Some(observers), "info", "t", &format!("m{i}"), None);
}
}
#[test]
fn recent_logs_after_answers_only_newer_entries() {
let observers = WorkerObservers::default();
push_messages(&observers, 3);
let (all, newest) = recent_logs_after(&observers, 0);
assert_eq!((all.len(), newest), (3, 3));
let (newer, _) = recent_logs_after(&observers, 2);
assert_eq!(newer.len(), 1);
assert_eq!(newer[0].message, "m2");
assert!(recent_logs_after(&observers, 3).0.is_empty());
}
#[test]
fn the_daemons_observers_share_the_global_worker_log() {
let a = WorkerObservers::with_global_worker_log();
let b = WorkerObservers::with_global_worker_log();
assert!(Arc::ptr_eq(&a.recent_logs, &b.recent_logs));
assert!(!Arc::ptr_eq(
&a.recent_logs,
&WorkerObservers::default().recent_logs
));
}
#[test]
fn recent_logs_after_a_restart_answers_the_whole_ring() {
let observers = WorkerObservers::default();
push_messages(&observers, 2);
assert_eq!(recent_logs_after(&observers, 99).0.len(), 2);
}
#[test]
fn recent_logs_after_skips_entries_that_left_the_ring() {
let observers = WorkerObservers::default();
push_messages(&observers, RECENT_LOGS_CAP + 5);
let (all, newest) = recent_logs_after(&observers, 0);
assert_eq!(all.len(), RECENT_LOGS_CAP);
assert_eq!(all[0].message, "m5");
assert_eq!(newest, (RECENT_LOGS_CAP + 5) as u64);
}
#[test]
fn is_unsupported_kind_detects_typed_unsupported_task() {
let err: anyhow::Error =
crate::engine::UnsupportedTask::new("synthetic", TaskKind::Llm).into();
assert!(is_unsupported_kind(&err));
assert!(err.to_string().contains("cannot serve llm"));
}
#[test]
fn is_unsupported_kind_survives_context_wrapping() {
let err = anyhow::Error::from(crate::engine::UnsupportedTask::new(
"sdcpp",
TaskKind::AudioTts,
))
.context("dispatching job j-1");
assert!(is_unsupported_kind(&err));
}
fn entry(message: &str) -> LogEntry {
LogEntry {
ts: Utc::now().to_rfc3339_opts(SecondsFormat::Millis, true),
level: "info".into(),
category: "test".into(),
message: message.into(),
job_id: None,
}
}
#[test]
fn restore_unshipped_requeues_batch_ahead_of_newer_entries() {
let logs: Arc<Mutex<Vec<LogEntry>>> = Arc::new(Mutex::new(vec![entry("newer")]));
restore_unshipped(&logs, vec![entry("batch-1"), entry("batch-2")]);
let queue = logs.lock();
let order: Vec<&str> = queue.iter().map(|e| e.message.as_str()).collect();
assert_eq!(order, vec!["batch-1", "batch-2", "newer"]);
}
#[test]
fn restore_unshipped_respects_the_queue_cap() {
let logs: Arc<Mutex<Vec<LogEntry>>> =
Arc::new(Mutex::new(vec![entry("newest"); LOG_SHIP_QUEUE_CAP]));
restore_unshipped(&logs, vec![entry("old-batch"); 100]);
let queue = logs.lock();
assert_eq!(queue.len(), LOG_SHIP_QUEUE_CAP);
assert_eq!(
queue.last().map(|e| e.message.as_str()),
Some("newest"),
"newest entries must survive the cap"
);
}
#[test]
fn ship_queue_is_bounded_and_records_dropped_entries() {
let logs: Arc<Mutex<Vec<LogEntry>>> = Arc::new(Mutex::new(Vec::new()));
for i in 0..(LOG_SHIP_QUEUE_CAP + 100) {
push_log_with_observers(&logs, None, "info", "test", &format!("entry {i}"), None);
}
let queue = logs.lock();
assert!(
queue.len() <= LOG_SHIP_QUEUE_CAP,
"ship queue exceeded its cap: {}",
queue.len()
);
assert_eq!(
queue.last().map(|e| e.message.as_str()),
Some(format!("entry {}", LOG_SHIP_QUEUE_CAP + 99).as_str())
);
assert!(
queue
.iter()
.any(|e| e.level == "warn" && e.message.contains("dropped")),
"overflow must leave a visible drop marker"
);
}
#[test]
fn recent_logs_ring_is_bounded_at_recent_logs_cap() {
let logs: Arc<Mutex<Vec<LogEntry>>> = Arc::new(Mutex::new(Vec::new()));
let observers = WorkerObservers::default();
let overflow = 25;
for i in 0..(RECENT_LOGS_CAP + overflow) {
push_log_with_observers(
&logs,
Some(&observers),
"info",
"test",
&format!("entry {i}"),
None,
);
}
let ring = observers.recent_logs.lock();
assert_eq!(
ring.len(),
RECENT_LOGS_CAP,
"the recent-logs ring must cap at RECENT_LOGS_CAP"
);
assert_eq!(
ring.back().map(|e| e.message.as_str()),
Some(format!("entry {}", RECENT_LOGS_CAP + overflow - 1).as_str()),
"the newest entry must survive at the back of the ring"
);
assert_eq!(
ring.front().map(|e| e.message.as_str()),
Some(format!("entry {overflow}").as_str()),
"the oldest surviving entry must be entry #overflow (older evicted)"
);
}
#[test]
fn session_state_summaries_carry_recovery_actions_for_terminal_states() {
assert!(SessionState::default() == SessionState::WaitingForApproval);
assert!(SessionState::Connected.summary().contains("connected"));
assert!(SessionState::Reconnecting { attempt: 3 }
.summary()
.contains("attempt 3"));
let auth = SessionState::AuthFailed {
reason: "bad token".into(),
}
.summary();
assert!(auth.contains("register --reset"), "got: {auth}");
assert!(auth.contains("bad token"));
assert!(SessionState::Fatal {
reason: "boom".into()
}
.summary()
.contains("boom"));
}
#[test]
fn set_session_state_updates_the_observer_slot() {
let observers = WorkerObservers::default();
assert_eq!(
*observers.session_state.lock(),
SessionState::WaitingForApproval
);
set_session_state(&observers, SessionState::Connected);
assert_eq!(*observers.session_state.lock(), SessionState::Connected);
}
#[test]
fn sync_studio_model_mirrors_into_the_shared_catalog_and_persists() {
use crate::types::{ModelCliDefaults, ModelEngine, ModelSource};
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("models.json");
let observers = WorkerObservers::default();
*observers.catalog_path.lock() = Some(path.clone());
let source = ModelSource {
engine: ModelEngine::Synthetic,
files: vec![],
cli_defaults: ModelCliDefaults::default(),
};
sync_studio_model(&observers, "studio-llm", TaskKind::Llm, &source);
assert!(observers.catalog.lock().get("studio-llm").is_some());
assert_eq!(
observers.catalog.lock().get("studio-llm").unwrap().origin,
"studio"
);
let reloaded = crate::catalog::Catalog::load_or_seed(&path).unwrap();
assert!(reloaded.get("studio-llm").is_some());
}
#[test]
fn sync_studio_model_without_a_path_stays_in_memory_only() {
use crate::types::{ModelCliDefaults, ModelEngine, ModelSource};
let observers = WorkerObservers::default();
let source = ModelSource {
engine: ModelEngine::Synthetic,
files: vec![],
cli_defaults: ModelCliDefaults::default(),
};
sync_studio_model(&observers, "m", TaskKind::Image, &source);
assert!(observers.catalog.lock().get("m").is_some());
}
#[test]
fn set_gpu_runtime_status_records_ok_without_warning() {
let observers = WorkerObservers::default();
assert!(observers.gpu_runtime.lock().is_none(), "unprobed at first");
let out = crate::test_support::capture({
let observers = observers.clone();
move || set_gpu_runtime_status(&observers, Ok(()))
});
let status = observers.gpu_runtime.lock().clone().unwrap();
assert!(status.ok);
assert!(status.detail.contains("available"));
assert!(
!out.contains("GPU runtime missing"),
"the ok path must not warn: {out}"
);
}
#[test]
fn set_gpu_runtime_status_records_and_warns_the_remedy_when_missing() {
let observers = WorkerObservers::default();
let out = crate::test_support::capture({
let observers = observers.clone();
move || {
set_gpu_runtime_status(
&observers,
Err(anyhow!("Vulkan runtime not available: install libvulkan1")),
)
}
});
let status = observers.gpu_runtime.lock().clone().unwrap();
assert!(!status.ok);
assert!(
status.detail.contains("libvulkan1"),
"got: {}",
status.detail
);
assert!(
out.contains("GPU runtime missing") && out.contains("WARN"),
"a missing runtime must warn with the remedy: {out}"
);
}
#[test]
fn capabilities_advertises_all_synthetic_kinds() {
let cfg = Config::default();
let engine = SyntheticEngine::new();
let cap = build_capabilities(&cfg, &engine);
assert_eq!(cap.engine, "synthetic");
assert_eq!(cap.task_kinds.len(), TaskKind::ALL.len());
assert!(cap.auto_enabled, "default capability snapshot is unpaused");
for kind in TaskKind::ALL {
assert!(cap.supported_models_per_kind.contains_key(&kind));
}
}
#[test]
fn capabilities_with_paused_flag_drives_auto_enabled() {
let cfg = Config::default();
let engine = SyntheticEngine::new();
let paused_caps = build_capabilities_with(&cfg, &engine, false);
assert!(!paused_caps.auto_enabled);
}
#[test]
fn summarize_capabilities_lists_engine_kinds_models_vram_and_pause_state() {
let cfg = Config {
vram_threshold_gb: 6.0,
..Config::default()
};
let engine = SyntheticEngine::new();
let caps = build_capabilities_with(&cfg, &engine, true);
let summary = summarize_capabilities(&caps);
assert!(summary.contains("engine=synthetic"), "got: {summary}");
for kind in &caps.task_kinds {
assert!(
summary.contains(kind.as_str()),
"missing kind {} in: {summary}",
kind.as_str()
);
}
assert!(
summary.contains(&format!("{} model(s)", caps.supported_models.len())),
"missing model count in: {summary}"
);
assert!(
summary.contains("synthetic"),
"missing model id in: {summary}"
);
assert!(
summary.contains("6.0"),
"missing vram threshold in: {summary}"
);
assert!(summary.contains("auto_enabled=true"), "got: {summary}");
}
#[test]
fn summarize_capabilities_reflects_paused_state() {
let cfg = Config::default();
let engine = SyntheticEngine::new();
let caps = build_capabilities_with(&cfg, &engine, false);
assert!(
summarize_capabilities(&caps).contains("auto_enabled=false"),
"paused worker must advertise auto_enabled=false"
);
}
fn caps_with_vram(total_gb: f32, threshold_gb: f32) -> WorkerCapabilities {
let mut caps = build_capabilities_with(&Config::default(), &SyntheticEngine::new(), true);
caps.vram_total_gb = total_gb;
caps.vram_threshold_gb = threshold_gb;
caps
}
#[test]
fn vram_threshold_warning_flags_threshold_above_detected_vram() {
let warning = vram_threshold_warning(&caps_with_vram(8.0, 12.0))
.expect("threshold above detected VRAM must warn");
assert!(warning.contains("12.0"), "missing threshold in: {warning}");
assert!(
warning.contains("8.0"),
"missing detected VRAM in: {warning}"
);
assert!(
warning.contains("vram_threshold_gb"),
"must name the config key to change: {warning}"
);
}
#[test]
fn vram_threshold_warning_silent_when_threshold_within_detected_vram() {
assert!(vram_threshold_warning(&caps_with_vram(24.0, 12.0)).is_none());
}
#[test]
fn vram_threshold_warning_silent_when_threshold_equals_detected() {
assert!(vram_threshold_warning(&caps_with_vram(12.0, 12.0)).is_none());
}
#[test]
fn vram_threshold_warning_silent_when_vram_undetected() {
assert!(vram_threshold_warning(&caps_with_vram(0.0, 12.0)).is_none());
}
#[test]
fn prompt_for_extracts_per_kind() {
let image = Task::Image(ImageParams {
prompt: "a stone golem".into(),
..Default::default()
});
assert_eq!(prompt_for(&image), "a stone golem");
let llm = Task::Llm(LlmParams {
messages: vec![
ChatMessage {
role: "system".into(),
content: "be helpful".into(),
},
ChatMessage {
role: "user".into(),
content: "hi".into(),
},
],
max_tokens: 32,
temperature: 0.5,
..Default::default()
});
assert_eq!(prompt_for(&llm), "hi");
let llm_empty = Task::Llm(LlmParams {
messages: vec![],
..Default::default()
});
assert_eq!(prompt_for(&llm_empty), "");
let stt = Task::AudioStt(AudioSttParams {
input_url: "https://example.com/clip.wav".into(),
..Default::default()
});
assert_eq!(prompt_for(&stt), "https://example.com/clip.wav");
let tts = Task::AudioTts(AudioTtsParams {
text: "hi there".into(),
voice: "v".into(),
ext: "wav".into(),
..Default::default()
});
assert_eq!(prompt_for(&tts), "hi there");
let video = Task::Video(VideoParams {
prompt: "a tiny dragon".into(),
seconds: 1.0,
width: 256,
height: 256,
ext: "mp4".into(),
..Default::default()
});
assert_eq!(prompt_for(&video), "a tiny dragon");
}
#[test]
fn truncate_prompt_passes_short_through_and_clips_long_prompts() {
let short = "a stone golem";
assert_eq!(truncate_prompt(short), short);
let exactly = "x".repeat(PROMPT_PREVIEW_CHARS);
assert_eq!(
truncate_prompt(&exactly),
exactly,
"a prompt exactly at the cap must not be clipped"
);
let over = "y".repeat(PROMPT_PREVIEW_CHARS + 1);
let clipped = truncate_prompt(&over);
assert_eq!(
clipped.chars().count(),
PROMPT_PREVIEW_CHARS + 1,
"clipped preview is the cap plus one ellipsis char"
);
assert!(
clipped.ends_with('\u{2026}'),
"a clipped preview ends with an ellipsis"
);
assert_eq!(
clipped
.chars()
.take(PROMPT_PREVIEW_CHARS)
.collect::<String>(),
"y".repeat(PROMPT_PREVIEW_CHARS),
"the kept prefix is the first PROMPT_PREVIEW_CHARS chars"
);
}
#[test]
fn truncate_prompt_clips_on_char_boundaries_for_multibyte_text() {
let multibyte = "\u{3042}".repeat(PROMPT_PREVIEW_CHARS + 1);
let clipped = truncate_prompt(&multibyte);
assert_eq!(clipped.chars().count(), PROMPT_PREVIEW_CHARS + 1);
assert!(clipped.ends_with('\u{2026}'));
assert_eq!(
clipped.chars().filter(|c| *c == '\u{3042}').count(),
PROMPT_PREVIEW_CHARS,
"exactly PROMPT_PREVIEW_CHARS multibyte chars survive the clip"
);
}
#[test]
fn is_unsupported_kind_matches_engine_message() {
let err = anyhow!("multi engine cannot serve llm tasks");
assert!(is_unsupported_kind(&err));
let other = anyhow!("network timeout");
assert!(!is_unsupported_kind(&other));
}
#[test]
fn format_status_includes_every_field() {
let cfg = Config::default();
let out = format_status(&cfg, std::path::Path::new("/tmp/x.toml"));
assert!(out.contains("config path:"));
assert!(out.contains("api_base_url:"));
assert!(out.contains("registration:"));
assert!(out.contains("not registered"));
assert!(out.contains("models_root:"));
assert!(out.contains("auto_update:"));
assert!(out.contains("update_interval:"));
}
#[test]
fn format_status_shows_worker_id_when_registered() {
let cfg = Config {
worker_id: Some("w-abc".into()),
auth_token: Some("tok".into()),
..Config::default()
};
let out = format_status(&cfg, std::path::Path::new("/tmp/x.toml"));
assert!(out.contains("w-abc"));
assert!(out.contains("approved"));
}
#[test]
fn format_status_shows_pending_request_id() {
let cfg = Config {
registration_request_id: Some("rr-7".into()),
..Config::default()
};
let out = format_status(&cfg, std::path::Path::new("/tmp/x.toml"));
assert!(out.contains("pending operator approval"));
assert!(out.contains("rr-7"));
}
#[test]
fn format_check_outcome_handles_both_branches() {
let up = update::CheckOutcome::UpToDate {
current: semver::Version::new(1, 2, 3),
};
assert!(format_check_outcome(&up).contains("up to date"));
let newer = update::CheckOutcome::NewerAvailable {
current: semver::Version::new(1, 2, 3),
latest: semver::Version::new(1, 3, 0),
};
let s = format_check_outcome(&newer);
assert!(s.contains("1.2.3 -> 1.3.0"));
}
#[test]
fn push_log_appends_an_entry() {
let logs: Arc<Mutex<Vec<LogEntry>>> = Arc::new(Mutex::new(Vec::new()));
push_log(&logs, "info", "test", "hi", None);
push_log(&logs, "warn", "test", "wat", Some("j-1".into()));
push_log(&logs, "error", "test", "boom", None);
let v = logs.lock();
assert_eq!(v.len(), 3);
assert_eq!(v[0].level, "info");
assert_eq!(v[1].level, "warn");
assert_eq!(v[1].job_id.as_deref(), Some("j-1"));
assert_eq!(v[2].level, "error");
}
#[test]
fn push_log_emits_job_id_as_a_structured_tracing_field() {
use crate::test_support::capture;
let logs = capture(|| {
let logs: Arc<Mutex<Vec<LogEntry>>> = Arc::new(Mutex::new(Vec::new()));
push_log(
&logs,
"info",
"ws",
"binary upload ok",
Some("job-42".into()),
);
});
assert!(
logs.contains("job_id=\"job-42\""),
"expected structured job_id field, got: {logs}"
);
assert!(
logs.contains("[ws] binary upload ok"),
"expected the human-readable message to survive, got: {logs}"
);
}
#[test]
fn push_log_omits_job_id_field_when_absent() {
use crate::test_support::capture;
let logs = capture(|| {
let logs: Arc<Mutex<Vec<LogEntry>>> = Arc::new(Mutex::new(Vec::new()));
push_log(&logs, "info", "auto-update", "up to date", None);
});
assert!(
!logs.contains("job_id"),
"expected no job_id field for a jobless log, got: {logs}"
);
}
#[test]
fn request_shutdown_sets_the_stop_flag() {
let stop = AtomicBool::new(false);
request_shutdown(&stop, "SIGTERM");
assert!(stop.load(Ordering::SeqCst));
}
#[test]
fn request_shutdown_reconfirms_when_already_stopping() {
let stop = AtomicBool::new(true);
request_shutdown(&stop, "SIGINT");
assert!(stop.load(Ordering::SeqCst));
}
#[test]
fn request_shutdown_emits_a_named_shutdown_breadcrumb() {
use crate::test_support::capture;
let logs = capture(|| {
let stop = AtomicBool::new(false);
request_shutdown(&stop, "SIGTERM");
});
assert!(logs.contains("INFO"), "expected INFO event, got: {logs}");
assert!(
logs.contains("studio_worker::runtime"),
"expected runtime target, got: {logs}"
);
assert!(
logs.contains("op=\"shutdown\""),
"expected op field, got: {logs}"
);
assert!(
logs.contains("signal=\"SIGTERM\""),
"expected signal field, got: {logs}"
);
}
#[tokio::test]
async fn auto_update_tick_disabled_when_flag_off() {
let cfg = Config {
auto_update_enabled: false,
..Config::default()
};
let logs = Arc::new(Mutex::new(Vec::new()));
let decision = auto_update_tick(&cfg, &crate::job_gate::JobGate::new(), &logs).await;
assert_eq!(decision, AutoUpdateDecision::Disabled);
}
#[tokio::test]
async fn auto_update_tick_skipped_when_busy() {
let cfg = Config {
auto_update_enabled: true,
..Config::default()
};
let logs = Arc::new(Mutex::new(Vec::new()));
let gate = crate::job_gate::JobGate::new();
let _held = gate.try_reserve().expect("hold the slot");
let decision = auto_update_tick(&cfg, &gate, &logs).await;
assert_eq!(decision, AutoUpdateDecision::SkippedBusy);
let entries = logs.lock();
assert!(entries.iter().any(|e| e.message.contains("busy on a job")));
}
#[tokio::test]
async fn wait_with_stop_short_circuits_when_already_stopped() {
let stop = Arc::new(AtomicBool::new(true));
let start = std::time::Instant::now();
wait_with_stop(Duration::from_secs(60), &stop, Duration::from_millis(10)).await;
assert!(
start.elapsed() < Duration::from_millis(100),
"an already-set stop must return without sleeping the full duration"
);
}
#[tokio::test]
async fn auto_updater_stops_promptly_during_idle_wait() {
let cfg = crate::config::shared(Config {
auto_update_enabled: false,
..Config::default()
});
let stop = Arc::new(AtomicBool::new(false));
let logs: Arc<Mutex<Vec<LogEntry>>> = Arc::new(Mutex::new(Vec::new()));
let busy = Arc::new(AtomicBool::new(false));
let schedule = LoopSchedule {
ws_session: crate::ws::session::SessionSchedule::fast_for_tests(),
auto_update_tick: Duration::from_secs(3600),
shutdown_tick: Duration::from_millis(1),
};
let handle = spawn_auto_updater(cfg, stop.clone(), logs, busy, schedule);
tokio::time::sleep(Duration::from_millis(10)).await;
stop.store(true, Ordering::SeqCst);
tokio::time::timeout(Duration::from_millis(250), handle)
.await
.expect("auto-updater did not observe stop promptly")
.expect("auto-updater task panicked");
}
#[test]
fn resolve_local_api_port_prefers_env_then_config_then_default() {
assert_eq!(resolve_local_api_port(Some("5000"), Some(4000)), 5000);
assert_eq!(resolve_local_api_port(None, Some(4000)), 4000);
assert_eq!(resolve_local_api_port(None, None), DEFAULT_LOCAL_API_PORT);
}
#[test]
fn resolve_local_api_port_warns_on_invalid_env_and_falls_back() {
let logs = crate::test_support::capture(|| {
assert_eq!(
resolve_local_api_port(Some("not-a-port"), None),
DEFAULT_LOCAL_API_PORT
);
assert_eq!(resolve_local_api_port(Some("99999"), Some(4001)), 4001);
});
assert!(logs.contains("WARN"), "expected a WARN, got: {logs}");
assert!(
logs.contains("not-a-port"),
"the warn must name the invalid value, got: {logs}"
);
assert!(
logs.contains("STUDIO_WORKER_LOCAL_API_PORT"),
"the warn must name the env var, got: {logs}"
);
}
#[test]
fn ensure_local_api_token_mints_once_and_persists() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("config.toml");
let cfg = config::shared(Config::default());
let minted = ensure_local_api_token(&cfg, &path);
assert_eq!(minted.len(), 64, "expected a 64-hex token");
assert!(minted.chars().all(|c| c.is_ascii_hexdigit()));
let (loaded, _) = config::load(Some(&path.to_string_lossy())).unwrap();
assert_eq!(loaded.local_api_token.as_deref(), Some(minted.as_str()));
let again = ensure_local_api_token(&cfg, &path);
assert_eq!(again, minted);
}
#[test]
fn ensure_local_api_token_keeps_an_existing_token() {
let dir = tempfile::tempdir().unwrap();
let path = dir.path().join("config.toml");
let cfg = config::shared(Config {
local_api_token: Some("pre-existing".into()),
..Config::default()
});
assert_eq!(ensure_local_api_token(&cfg, &path), "pre-existing");
assert!(!path.exists(), "no save when nothing changed");
}
#[test]
fn ensure_local_api_token_survives_a_failed_persist() {
let dir = tempfile::tempdir().unwrap();
let blocked = dir.path().join("blocked");
std::fs::write(&blocked, b"a file, not a dir").unwrap();
let path = blocked.join("config.toml");
let cfg = config::shared(Config::default());
let logs = crate::test_support::capture({
let cfg = cfg.clone();
move || {
let token = ensure_local_api_token(&cfg, &path);
assert_eq!(token.len(), 64);
}
});
assert!(
logs.contains("failed to persist the local api token"),
"a failed persist must warn: {logs}"
);
assert!(
cfg.lock().local_api_token.is_some(),
"the in-memory token must survive the failed persist"
);
}
}