use std::sync::{Arc, RwLock};
use axum::extract::State;
use axum::http::StatusCode;
use axum::response::{IntoResponse, Response};
use axum::Json;
use frink_api::{LoraAdapterInfo, LoraApplyResponse, LoraScaleRequest};
use frink_gguf::TensorSource;
use frink_models::lora_attach::LoraSpec;
use frink_models::Decoder;
use crate::{invalid_request, unsupported_feature, ApiError, AppState, Model};
pub(crate) const ENV_SPECS: &str = "FRINK_LORA";
pub(crate) const ENV_INIT_WITHOUT_APPLY: &str = "FRINK_LORA_INIT_WITHOUT_APPLY";
pub(crate) fn specs_from_env() -> anyhow::Result<Vec<LoraSpec>> {
match std::env::var(ENV_SPECS) {
Ok(raw) => parse_specs(&raw),
Err(_) => Ok(Vec::new()),
}
}
fn parse_specs(raw: &str) -> anyhow::Result<Vec<LoraSpec>> {
raw.split(',')
.map(str::trim)
.filter(|s| !s.is_empty())
.map(|s| LoraSpec::parse_scaled(s).map_err(|e| anyhow::anyhow!("{ENV_SPECS}: {e}")))
.collect()
}
fn init_without_apply() -> bool {
std::env::var(ENV_INIT_WITHOUT_APPLY).is_ok_and(|v| v == "1")
}
pub(crate) fn attach_from_env(
decoder: &mut Decoder,
base: &impl TensorSource,
) -> anyhow::Result<()> {
let specs = specs_from_env()?;
if specs.is_empty() {
return Ok(());
}
decoder
.attach_lora_specs(base, &specs)
.map_err(|e| anyhow::anyhow!("lora: {e}"))?;
if init_without_apply() {
decoder
.set_lora_scales(&[])
.map_err(|e| anyhow::anyhow!("lora: {e}"))?;
tracing::info!(
"{} adapter(s) loaded at scale 0 ({ENV_INIT_WITHOUT_APPLY}); apply them with \
POST /lora-adapters",
specs.len()
);
}
Ok(())
}
pub(crate) fn refuse_env_for_engine(engine: &str) -> anyhow::Result<()> {
if specs_from_env()?.is_empty() {
return Ok(());
}
anyhow::bail!(
"{ENV_SPECS} (--lora) is not implemented for the {engine} engine: only the generic \
decoder attaches adapters; refusing rather than serving the base weights"
)
}
static GATE: RwLock<()> = RwLock::new(());
enum Guard {
Read(#[allow(dead_code)] std::sync::RwLockReadGuard<'static, ()>),
Write(#[allow(dead_code)] std::sync::RwLockWriteGuard<'static, ()>),
}
pub(crate) struct LoraLease {
_guard: Guard,
restore: Option<(Arc<Decoder>, Vec<f32>)>,
}
impl Drop for LoraLease {
fn drop(&mut self) {
if let Some((decoder, previous)) = self.restore.take() {
let scales: Vec<(usize, f32)> = previous.into_iter().enumerate().collect();
let _ = decoder.set_lora_scales(&scales);
}
}
}
pub(crate) fn lease(model: &Model, scales: Option<&[f32]>) -> LoraLease {
let shared = || LoraLease {
_guard: Guard::Read(GATE.read().unwrap_or_else(|e| e.into_inner())),
restore: None,
};
let Some(want) = scales else {
return shared();
};
let Model::Gguf(g) = model else {
return shared();
};
let read = GATE.read().unwrap_or_else(|e| e.into_inner());
if g.decoder.lora_scales() == want {
return LoraLease {
_guard: Guard::Read(read),
restore: None,
};
}
drop(read);
let guard = GATE.write().unwrap_or_else(|e| e.into_inner());
let previous = g.decoder.lora_scales();
let scales: Vec<(usize, f32)> = want.iter().copied().enumerate().collect();
let _ = g.decoder.set_lora_scales(&scales);
LoraLease {
_guard: Guard::Write(guard),
restore: Some((Arc::clone(&g.decoder), previous)),
}
}
pub(crate) fn resolve_request(
model: &Model,
requested: Option<&[LoraScaleRequest]>,
) -> Result<Option<Vec<f32>>, ApiError> {
let Model::Gguf(g) = model else {
return match requested {
Some(_) => Err(unsupported_feature(
"`lora` is only served by the generic decoder; this checkpoint runs on a \
dedicated engine that attaches no adapters",
)),
None => Ok(None),
};
};
let n = g.decoder.lora_adapters.len();
let current = || (n > 0).then(|| g.decoder.lora_scales());
let Some(requested) = requested else {
return Ok(current());
};
if requested.is_empty() {
return Ok(current());
}
if n == 0 {
return Err(invalid_request(
"`lora` names an adapter but none is loaded (start the server with --lora)",
"lora",
));
}
let mut want = vec![0f32; n];
for entry in requested {
if entry.id >= n {
return Err(invalid_request(
&format!(
"lora adapter id {} is out of range: {n} adapter(s) loaded",
entry.id
),
"lora",
));
}
want[entry.id] = entry.scale;
}
Ok(Some(want))
}
pub(crate) fn list(model: &Model) -> Vec<LoraAdapterInfo> {
let Model::Gguf(g) = model else {
return Vec::new();
};
g.decoder
.lora_adapters
.iter()
.enumerate()
.map(|(id, a)| LoraAdapterInfo {
id,
path: a.path.display().to_string(),
scale: a.scale(),
task_name: a.task_name.clone(),
prompt_prefix: a.prompt_prefix.clone(),
})
.collect()
}
pub(crate) async fn get_lora_adapters(State(state): State<Arc<AppState>>) -> Response {
let active = match state.require_active() {
Ok(a) => a,
Err(e) => return e.into_response(),
};
let model = match active.generative() {
Ok(m) => m,
Err(_) => return Json(Vec::<LoraAdapterInfo>::new()).into_response(),
};
Json(list(model)).into_response()
}
pub(crate) async fn post_lora_adapters(
State(state): State<Arc<AppState>>,
Json(body): Json<Vec<LoraScaleRequest>>,
) -> Response {
let active = match state.require_active() {
Ok(a) => a,
Err(e) => return e.into_response(),
};
let model = match active.generative() {
Ok(m) => Arc::clone(m),
Err(e) => return e.into_response(),
};
let Model::Gguf(_) = &*model else {
return unsupported_feature(
"this checkpoint runs on a dedicated engine that attaches no adapters",
)
.into_response();
};
let scales: Vec<(usize, f32)> = body.iter().map(|e| (e.id, e.scale)).collect();
let result = tokio::task::spawn_blocking(move || {
let _gate = GATE.write().unwrap_or_else(|e| e.into_inner());
let Model::Gguf(g) = &*model else {
unreachable!("checked above");
};
g.decoder.set_lora_scales(&scales)
})
.await;
match result {
Ok(Ok(())) => Json(LoraApplyResponse { success: true }).into_response(),
Ok(Err(msg)) => invalid_request(&msg, "id").into_response(),
Err(e) => (
StatusCode::INTERNAL_SERVER_ERROR,
Json(serde_json::json!({"error": {"message": format!("lora apply task failed: {e}")}})),
)
.into_response(),
}
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn env_specs_parse_in_order_and_refuse_a_bad_scale() {
let specs = parse_specs("a.gguf:1, b.gguf:0.5 ,").unwrap();
assert_eq!(specs.len(), 2);
assert_eq!(specs[1].scale, 0.5);
assert!(parse_specs("a.gguf")
.unwrap_err()
.to_string()
.contains("FNAME:SCALE"));
assert!(parse_specs("").unwrap().is_empty());
}
#[test]
fn a_lease_without_an_override_is_shared_and_restores_nothing() {
let a = LoraLease {
_guard: Guard::Read(GATE.read().unwrap()),
restore: None,
};
let b = LoraLease {
_guard: Guard::Read(GATE.read().unwrap()),
restore: None,
};
assert!(GATE.try_write().is_err());
drop(a);
drop(b);
drop(GATE.write().unwrap());
}
}
#[cfg(test)]
mod http_tests {
use std::sync::Arc;
use std::time::Duration;
use axum::http::StatusCode;
use frink_core::weight_matrix::{LoraDelta, LoraScale};
use frink_models::lora_attach::LoraAttached;
use frink_models::Decoder;
use crate::response_cache::ResponseCache;
use crate::tests::{get_json, post_json_uri, test_app_with_state, test_state};
use crate::{chat_template, GgufModel, Model, ServerTokenizer, StopTokens};
fn model_with_adapters(n: usize) -> (Model, Arc<Decoder>) {
let mut cfg = frink_models::config::test_dense_fixture();
cfg.vocab_size = 256;
let mut d = Decoder::new_random_small(cfg, 2, 256);
for i in 0..n {
let head = &mut d.output_head;
let (rows, cols) = (head.rows(), head.cols());
let scale = LoraScale::new(1.0);
let mut b = vec![0.0; rows];
b[10 + i] = 4.0;
head.attach_lora(
LoraDelta::new(vec![1.0; cols], b, 1, rows, cols, 0.0, scale.clone()).unwrap(),
);
d.lora_adapters.push(LoraAttached {
path: format!("adapter_{i}.gguf").into(),
alpha: 0.0,
task_name: String::new(),
prompt_prefix: String::new(),
scale,
n_tensors: 1,
});
}
let decoder = Arc::new(d);
let model = Model::Gguf(GgufModel {
decoder: Arc::clone(&decoder),
tokenizer: Arc::new(ServerTokenizer::Byte),
stop_tokens: StopTokens::default(),
bos_id: None,
is_synthetic: true,
chat_template: chat_template::PromptTemplate::plain(),
});
(model, decoder)
}
fn app(model: Model) -> axum::Router {
test_app_with_state(Arc::new(test_state(
model,
ResponseCache::new(16, Duration::from_secs(60)),
)))
}
#[tokio::test]
async fn a_model_without_adapters_lists_none_and_refuses_the_field_by_name() {
let (model, _) = model_with_adapters(0);
let app = app(model);
let (status, body) = get_json(&app, frink_api::routes::LORA_ADAPTERS).await;
assert_eq!(status, StatusCode::OK);
assert_eq!(body, serde_json::json!([]));
let (status, body) = post_json_uri(
&app,
frink_api::routes::COMPLETION,
serde_json::json!({"prompt": "hi", "n_predict": 1, "lora": [{"id": 0, "scale": 0.5}]}),
)
.await;
assert_eq!(status, StatusCode::BAD_REQUEST, "{body}");
assert!(
body["error"]["message"]
.as_str()
.unwrap()
.contains("none is loaded"),
"{body}"
);
let (status, _) = post_json_uri(
&app,
frink_api::routes::COMPLETION,
serde_json::json!({"prompt": "hi", "n_predict": 1, "lora": []}),
)
.await;
assert_eq!(status, StatusCode::OK);
}
#[tokio::test]
async fn get_lists_every_adapter_and_post_sets_the_scales_zeroing_the_unnamed() {
let (model, decoder) = model_with_adapters(2);
let app = app(model);
let (status, body) = get_json(&app, frink_api::routes::LORA_ADAPTERS).await;
assert_eq!(status, StatusCode::OK);
assert_eq!(
body,
serde_json::json!([
{"id": 0, "path": "adapter_0.gguf", "scale": 1.0, "task_name": "", "prompt_prefix": ""},
{"id": 1, "path": "adapter_1.gguf", "scale": 1.0, "task_name": "", "prompt_prefix": ""},
])
);
let (status, body) = post_json_uri(
&app,
frink_api::routes::LORA_ADAPTERS,
serde_json::json!([{"id": 1, "scale": 0.25}]),
)
.await;
assert_eq!(status, StatusCode::OK, "{body}");
assert_eq!(body, serde_json::json!({"success": true}));
assert_eq!(
decoder.lora_scales(),
vec![0.0, 0.25],
"unnamed adapter 0 went to 0"
);
let (_, body) = get_json(&app, frink_api::routes::LORA_ADAPTERS).await;
assert_eq!(body[0]["scale"], 0.0);
assert_eq!(body[1]["scale"], 0.25);
let (status, body) = post_json_uri(
&app,
frink_api::routes::LORA_ADAPTERS,
serde_json::json!([{"id": 2, "scale": 1.0}]),
)
.await;
assert_eq!(status, StatusCode::BAD_REQUEST, "{body}");
assert!(body["error"]["message"]
.as_str()
.unwrap()
.contains("out of range"));
assert_eq!(decoder.lora_scales(), vec![0.0, 0.25]);
let (status, _) = post_json_uri(
&app,
frink_api::routes::LORA_ADAPTERS,
serde_json::json!([{"id": 0}]),
)
.await;
assert_eq!(status, StatusCode::OK);
assert_eq!(decoder.lora_scales(), vec![0.0, 0.0]);
}
#[tokio::test]
async fn a_per_request_override_runs_alone_and_restores_the_scales_after() {
let (model, decoder) = model_with_adapters(2);
let app = app(model);
assert_eq!(decoder.lora_scales(), vec![1.0, 1.0]);
let run = |lora: serde_json::Value| {
let app = app.clone();
async move {
post_json_uri(
&app,
frink_api::routes::COMPLETION,
serde_json::json!({
"prompt": "abc", "n_predict": 3, "temperature": 0.0, "lora": lora
}),
)
.await
}
};
let (status, plain) =
run(serde_json::json!([{"id": 0, "scale": 1.0}, {"id": 1, "scale": 1.0}])).await;
assert_eq!(status, StatusCode::OK, "{plain}");
let content = |b: &serde_json::Value| b["content"].as_str().unwrap().to_string();
let mut hit = false;
for sign in [1000.0, -1000.0] {
let (status, overridden) = run(serde_json::json!([{"id": 1, "scale": sign}])).await;
assert_eq!(status, StatusCode::OK, "{overridden}");
assert_eq!(
decoder.lora_scales(),
vec![1.0, 1.0],
"the override lasted exactly one generation"
);
hit |= content(&overridden).contains("\\u{b}");
}
assert!(hit, "the override must have reached the weights: {plain}");
assert!(!content(&plain).contains("\\u{b}"));
let (status, body) = run(serde_json::json!([{"id": 5, "scale": 1.0}])).await;
assert_eq!(status, StatusCode::BAD_REQUEST, "{body}");
assert_eq!(decoder.lora_scales(), vec![1.0, 1.0]);
}
#[test]
fn a_lease_sets_the_scales_for_its_lifetime_and_the_head_sees_them() {
let (model, decoder) = model_with_adapters(2);
let mut kv = decoder.config.new_kv_caches();
let before = decoder.forward_token(3, 0, &mut kv);
{
let _lease = super::lease(&model, Some(&[0.0, 1000.0]));
assert_eq!(decoder.lora_scales(), vec![0.0, 1000.0]);
let mut kv = decoder.config.new_kv_caches();
let during = decoder.forward_token(3, 0, &mut kv);
let argmax = |v: &[f32]| {
v.iter()
.enumerate()
.max_by(|a, b| a.1.partial_cmp(b.1).unwrap())
.unwrap()
.0
};
assert_ne!(before, during, "the head must see the scale");
drop(_lease);
let _neg = super::lease(&model, Some(&[0.0, -1000.0]));
let mut kv = decoder.config.new_kv_caches();
let neg = decoder.forward_token(3, 0, &mut kv);
assert!(
argmax(&during) == 11 || argmax(&neg) == 11,
"one sign must make token 11 the greedy pick: {} / {}",
argmax(&during),
argmax(&neg)
);
}
assert_eq!(decoder.lora_scales(), vec![1.0, 1.0], "restored on drop");
}
#[tokio::test]
async fn the_openai_routes_take_the_same_field() {
let (model, decoder) = model_with_adapters(1);
let app = app(model);
let (status, body) = post_json_uri(
&app,
frink_api::routes::V1_COMPLETIONS,
serde_json::json!({"model": "m", "prompt": "hi", "max_tokens": 1, "lora": [{"id": 3}]}),
)
.await;
assert_eq!(status, StatusCode::BAD_REQUEST, "{body}");
let (status, body) = post_json_uri(
&app,
frink_api::routes::V1_CHAT_COMPLETIONS,
serde_json::json!({"model": "m", "messages": [{"role": "user", "content": "hi"}],
"max_tokens": 1, "lora": [{"id": 0, "scale": 0.5}]}),
)
.await;
assert_eq!(status, StatusCode::OK, "{body}");
assert_eq!(
decoder.lora_scales(),
vec![1.0],
"restored after the chat turn"
);
}
}