use std::path::PathBuf;
use std::sync::atomic::{AtomicBool, Ordering};
use std::sync::{Arc, Mutex};
use std::time::{Duration, Instant};
use async_trait::async_trait;
use crate::error::GatewayError;
use crate::local::error::LocalError;
use crate::local::server::{LaunchOptions, ServerGuard};
use crate::upstream::Upstream;
use crate::wire::{ChatRequest, ChatResponse};
const RESPAWN_COOLDOWN: Duration = Duration::from_secs(3);
struct LocalInner {
executable: PathBuf,
model_path: PathBuf,
options: LaunchOptions,
model_name: String,
guard: Mutex<ServerGuard>,
last_respawn: Mutex<Option<Instant>>,
shut_down: AtomicBool,
}
#[derive(Clone)]
pub(crate) struct LocalUpstream {
inner: Arc<LocalInner>,
http: reqwest::Client,
}
impl std::fmt::Debug for LocalUpstream {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.debug_struct("LocalUpstream")
.field("executable", &self.inner.executable)
.field("model_path", &self.inner.model_path)
.field("options", &self.inner.options)
.field("model_name", &self.inner.model_name)
.field("guard", &self.inner.guard)
.finish_non_exhaustive()
}
}
impl LocalUpstream {
#[must_use]
pub(crate) fn new(
guard: ServerGuard,
executable: PathBuf,
model_path: PathBuf,
options: LaunchOptions,
model_name: String,
) -> LocalUpstream {
LocalUpstream {
inner: Arc::new(LocalInner {
executable,
model_path,
options,
model_name,
guard: Mutex::new(guard),
last_respawn: Mutex::new(None),
shut_down: AtomicBool::new(false),
}),
http: crate::http_util::bounded_client(),
}
}
fn teardown(&self) -> Result<(), LocalError> {
self.inner.shut_down.store(true, Ordering::Release);
let mut guard = self
.inner
.guard
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
guard.shutdown()
}
fn recover_if_dead(inner: &LocalInner) -> Result<bool, LocalError> {
if inner.shut_down.load(Ordering::Acquire) {
return Ok(false);
}
let mut guard = inner
.guard
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if inner.shut_down.load(Ordering::Acquire) {
return Ok(false);
}
if guard.is_running()? {
return Ok(false);
}
let mut last = inner
.last_respawn
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
if let Some(at) = *last
&& at.elapsed() < RESPAWN_COOLDOWN
{
tracing::warn!(
model = %inner.model_name,
"llama-server dead but respawn cooldown active"
);
return Err(LocalError::RespawnCooldown {
model: inner.model_name.clone(),
});
}
tracing::warn!(
model = %inner.model_name,
port = guard.port(),
"llama-server exited after readiness; respawning"
);
match guard.respawn(
&inner.executable,
&inner.model_path,
&inner.options,
&inner.shut_down,
) {
Ok(()) => {
*last = Some(Instant::now());
tracing::info!(
model = %inner.model_name,
port = guard.port(),
"llama-server respawned"
);
Ok(true)
}
Err(error) => {
*last = Some(Instant::now());
tracing::error!(
model = %inner.model_name,
error = %error,
retryable = error.is_retryable(),
"llama-server respawn failed"
);
Err(error)
}
}
}
#[cfg(test)]
pub(crate) fn test_recover(&self) -> Result<bool, LocalError> {
Self::recover_if_dead(&self.inner)
}
async fn forward(
&self,
mut req: ChatRequest,
upstream_model: &str,
) -> Result<ChatResponse, GatewayError> {
let requested = std::mem::replace(&mut req.model, upstream_model.to_string());
let (base_url, api_key) = {
let guard = self
.inner
.guard
.lock()
.unwrap_or_else(std::sync::PoisonError::into_inner);
(guard.base_url(), guard.api_key().to_owned())
};
let mut builder = self
.http
.post(format!(
"{}/chat/completions",
base_url.trim_end_matches('/')
))
.json(&req);
if !api_key.is_empty() {
builder = builder.bearer_auth(&api_key);
}
let response = builder
.send()
.await
.map_err(GatewayError::upstream_transport)?;
let status = response.status();
if !status.is_success() {
let body =
crate::http_util::read_body_capped(response, crate::http_util::MAX_ERROR_BODY)
.await;
let body: String = body.chars().take(2000).collect();
return Err(GatewayError::UpstreamStatus {
status: status.as_u16(),
body,
});
}
let bytes = crate::http_util::read_bytes_capped(response, crate::http_util::MAX_JSON_BODY)
.await
.map_err(GatewayError::upstream_transport)?;
let mut parsed: ChatResponse =
serde_json::from_slice(&bytes).map_err(GatewayError::upstream_protocol)?;
parsed.model = requested;
Ok(parsed)
}
}
#[async_trait]
impl Upstream for LocalUpstream {
async fn send(
&self,
req: ChatRequest,
upstream_model: &str,
) -> Result<ChatResponse, GatewayError> {
match self.forward(req.clone(), upstream_model).await {
Ok(response) => Ok(response),
Err(error) if matches!(error, GatewayError::UpstreamTransport(_)) => {
let inner = Arc::clone(&self.inner);
let (tx, rx) = tokio::sync::oneshot::channel();
std::thread::spawn(move || {
let _ = tx.send(LocalUpstream::recover_if_dead(&inner));
});
match map_recovery_reply(rx.await, error) {
RecoveryOutcome::Retry => self.forward(req, upstream_model).await,
RecoveryOutcome::Failed(err) => Err(err),
}
}
Err(error) => Err(error),
}
}
fn shutdown(&self) -> Result<(), LocalError> {
self.teardown()
}
}
enum RecoveryOutcome {
Retry,
Failed(GatewayError),
}
fn map_recovery_reply(
reply: Result<Result<bool, LocalError>, tokio::sync::oneshot::error::RecvError>,
original: GatewayError,
) -> RecoveryOutcome {
match reply {
Ok(Ok(true)) => RecoveryOutcome::Retry,
Ok(Ok(false)) => RecoveryOutcome::Failed(original),
Ok(Err(local)) => RecoveryOutcome::Failed(GatewayError::UpstreamTransport(Box::new(local))),
Err(_) => RecoveryOutcome::Failed(GatewayError::UpstreamTransport(Box::new(
std::io::Error::other("llama-server respawn thread dropped before reporting"),
))),
}
}
#[cfg(test)]
mod tests {
use super::*;
fn transport_err() -> GatewayError {
GatewayError::UpstreamTransport(Box::new(std::io::Error::other(
"original transport failure",
)))
}
#[tokio::test]
async fn map_recovery_reply_covers_every_branch() {
assert!(matches!(
map_recovery_reply(Ok(Ok(true)), transport_err()),
RecoveryOutcome::Retry
));
assert!(matches!(
map_recovery_reply(Ok(Ok(false)), transport_err()),
RecoveryOutcome::Failed(GatewayError::UpstreamTransport(_))
));
assert!(matches!(
map_recovery_reply(Ok(Err(LocalError::TeardownTimeout)), transport_err()),
RecoveryOutcome::Failed(GatewayError::UpstreamTransport(_))
));
let (tx, rx) = tokio::sync::oneshot::channel::<std::result::Result<bool, LocalError>>();
drop(tx);
let dropped = rx.await;
match map_recovery_reply(dropped, transport_err()) {
RecoveryOutcome::Failed(GatewayError::UpstreamTransport(source)) => {
assert!(
source.to_string().contains("dropped before reporting"),
"unexpected message: {source}"
);
}
_ => panic!("dropped recovery reply must yield a transport failure"),
}
}
}