feat: add SSE streaming support (OpenAI-compatible)
- Token-by-token streaming via Server-Sent Events - tokio::sync::mpsc channel + ReceiverStream for clean async - OpenAI SSE format: role chunk → content chunks → finish chunk - Non-streaming still works with the same code path - Uses spawn_blocking pattern for CPU-bound inference
This commit is contained in:
Generated
+68
@@ -264,6 +264,21 @@ dependencies = [
|
|||||||
"percent-encoding",
|
"percent-encoding",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "futures"
|
||||||
|
version = "0.3.33"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "a88cf1f829d945f548cf8fec32c61b1f202b6d93b45848602fc02af4b12ad218"
|
||||||
|
dependencies = [
|
||||||
|
"futures-channel",
|
||||||
|
"futures-core",
|
||||||
|
"futures-executor",
|
||||||
|
"futures-io",
|
||||||
|
"futures-sink",
|
||||||
|
"futures-task",
|
||||||
|
"futures-util",
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "futures-channel"
|
name = "futures-channel"
|
||||||
version = "0.3.33"
|
version = "0.3.33"
|
||||||
@@ -271,6 +286,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index"
|
|||||||
checksum = "262590f4fe6afeb0bc83be1daa64e52657fe185690a958af7f3ad0e92085c5ae"
|
checksum = "262590f4fe6afeb0bc83be1daa64e52657fe185690a958af7f3ad0e92085c5ae"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
"futures-core",
|
"futures-core",
|
||||||
|
"futures-sink",
|
||||||
]
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
@@ -279,6 +295,40 @@ version = "0.3.33"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "2cd50c473c80f6d7c3670a752354b8e569b1a7cbfdc0419ec88e5edad85e0dc7"
|
checksum = "2cd50c473c80f6d7c3670a752354b8e569b1a7cbfdc0419ec88e5edad85e0dc7"
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "futures-executor"
|
||||||
|
version = "0.3.33"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "6754879cc9f2c66f88c6e5c35344bb0bdb0708b0352b1201815667c7eabc7458"
|
||||||
|
dependencies = [
|
||||||
|
"futures-core",
|
||||||
|
"futures-task",
|
||||||
|
"futures-util",
|
||||||
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "futures-io"
|
||||||
|
version = "0.3.33"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "4577ecaa3c4f96589d473f679a71b596316f6641bc350038b962a5daf0085d7a"
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "futures-macro"
|
||||||
|
version = "0.3.33"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "2d6d3cde68c518367be28956066ddfef33813991b77a55005a69dae04bf3b10b"
|
||||||
|
dependencies = [
|
||||||
|
"proc-macro2",
|
||||||
|
"quote",
|
||||||
|
"syn 2.0.119",
|
||||||
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "futures-sink"
|
||||||
|
version = "0.3.33"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "e34418ac499d6305c2fb5ad0ed2f6ac998c5f8ca209b4510f7f94242c647e307"
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "futures-task"
|
name = "futures-task"
|
||||||
version = "0.3.33"
|
version = "0.3.33"
|
||||||
@@ -291,8 +341,13 @@ version = "0.3.33"
|
|||||||
source = "registry+https://github.com/rust-lang/crates.io-index"
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
checksum = "a77a90a256fce34da66415271e30f94ee91c57b04b8a2c042d9cf3220179deaa"
|
checksum = "a77a90a256fce34da66415271e30f94ee91c57b04b8a2c042d9cf3220179deaa"
|
||||||
dependencies = [
|
dependencies = [
|
||||||
|
"futures-channel",
|
||||||
"futures-core",
|
"futures-core",
|
||||||
|
"futures-io",
|
||||||
|
"futures-macro",
|
||||||
|
"futures-sink",
|
||||||
"futures-task",
|
"futures-task",
|
||||||
|
"memchr",
|
||||||
"pin-project-lite",
|
"pin-project-lite",
|
||||||
"slab",
|
"slab",
|
||||||
]
|
]
|
||||||
@@ -511,10 +566,12 @@ dependencies = [
|
|||||||
"anyhow",
|
"anyhow",
|
||||||
"axum",
|
"axum",
|
||||||
"chrono",
|
"chrono",
|
||||||
|
"futures",
|
||||||
"llama-cpp-2",
|
"llama-cpp-2",
|
||||||
"serde",
|
"serde",
|
||||||
"serde_json",
|
"serde_json",
|
||||||
"tokio",
|
"tokio",
|
||||||
|
"tokio-stream",
|
||||||
"tower-http",
|
"tower-http",
|
||||||
"tracing",
|
"tracing",
|
||||||
"tracing-subscriber",
|
"tracing-subscriber",
|
||||||
@@ -958,6 +1015,17 @@ dependencies = [
|
|||||||
"syn 2.0.119",
|
"syn 2.0.119",
|
||||||
]
|
]
|
||||||
|
|
||||||
|
[[package]]
|
||||||
|
name = "tokio-stream"
|
||||||
|
version = "0.1.19"
|
||||||
|
source = "registry+https://github.com/rust-lang/crates.io-index"
|
||||||
|
checksum = "a3d06f0b082ba57c26b79407372e57cf2a1e28124f78e9479fe80322cf53420b"
|
||||||
|
dependencies = [
|
||||||
|
"futures-core",
|
||||||
|
"pin-project-lite",
|
||||||
|
"tokio",
|
||||||
|
]
|
||||||
|
|
||||||
[[package]]
|
[[package]]
|
||||||
name = "tower"
|
name = "tower"
|
||||||
version = "0.5.3"
|
version = "0.5.3"
|
||||||
|
|||||||
@@ -10,6 +10,7 @@ llama-cpp-2 = "0.1"
|
|||||||
# HTTP server
|
# HTTP server
|
||||||
axum = { version = "0.8", features = ["json"] }
|
axum = { version = "0.8", features = ["json"] }
|
||||||
tokio = { version = "1", features = ["full"] }
|
tokio = { version = "1", features = ["full"] }
|
||||||
|
tokio-stream = "0.1"
|
||||||
tower-http = { version = "0.6", features = ["cors", "trace"] }
|
tower-http = { version = "0.6", features = ["cors", "trace"] }
|
||||||
|
|
||||||
# Serialization
|
# Serialization
|
||||||
@@ -22,3 +23,4 @@ tracing-subscriber = { version = "0.3", features = ["env-filter"] }
|
|||||||
anyhow = "1"
|
anyhow = "1"
|
||||||
uuid = { version = "1", features = ["v4"] }
|
uuid = { version = "1", features = ["v4"] }
|
||||||
chrono = { version = "0.4", features = ["serde"] }
|
chrono = { version = "0.4", features = ["serde"] }
|
||||||
|
futures = "0.3"
|
||||||
|
|||||||
+187
-55
@@ -1,7 +1,10 @@
|
|||||||
use axum::{
|
use axum::{
|
||||||
extract::State,
|
extract::State,
|
||||||
http::{HeaderMap, StatusCode},
|
http::{HeaderMap, StatusCode},
|
||||||
response::Json,
|
response::{
|
||||||
|
sse::{Event, KeepAlive, Sse},
|
||||||
|
IntoResponse, Json,
|
||||||
|
},
|
||||||
routing::{get, post},
|
routing::{get, post},
|
||||||
Router,
|
Router,
|
||||||
};
|
};
|
||||||
@@ -15,9 +18,12 @@ use llama_cpp_2::{
|
|||||||
TokenToStringError,
|
TokenToStringError,
|
||||||
};
|
};
|
||||||
use serde::{Deserialize, Serialize};
|
use serde::{Deserialize, Serialize};
|
||||||
use std::num::NonZeroU32;
|
use std::{
|
||||||
use std::sync::Arc;
|
num::NonZeroU32,
|
||||||
|
sync::Arc,
|
||||||
|
};
|
||||||
use tokio::sync::Mutex;
|
use tokio::sync::Mutex;
|
||||||
|
use tokio_stream::wrappers::ReceiverStream;
|
||||||
use tower_http::cors::CorsLayer;
|
use tower_http::cors::CorsLayer;
|
||||||
use tracing::info;
|
use tracing::info;
|
||||||
|
|
||||||
@@ -27,7 +33,6 @@ struct CtxInner {
|
|||||||
sampler: LlamaSampler,
|
sampler: LlamaSampler,
|
||||||
}
|
}
|
||||||
|
|
||||||
// SAFETY: llama.cpp contexts are accessed from a single thread via the Mutex
|
|
||||||
unsafe impl Send for CtxInner {}
|
unsafe impl Send for CtxInner {}
|
||||||
unsafe impl Sync for CtxInner {}
|
unsafe impl Sync for CtxInner {}
|
||||||
|
|
||||||
@@ -45,10 +50,8 @@ impl CtxInner {
|
|||||||
self.context.decode(&mut batch).map_err(|e| e.to_string())
|
self.context.decode(&mut batch).map_err(|e| e.to_string())
|
||||||
}
|
}
|
||||||
|
|
||||||
// Use raw pointer to avoid borrow checker limitations with llama-cpp-2 API
|
|
||||||
fn sample_token(&mut self) -> LlamaToken {
|
fn sample_token(&mut self) -> LlamaToken {
|
||||||
let ctx_ptr: *const llama_cpp_2::context::LlamaContext = &self.context;
|
let ctx_ptr: *const llama_cpp_2::context::LlamaContext = &self.context;
|
||||||
// SAFETY: sampler is the sole owner of the context reference during this call
|
|
||||||
let ctx_ref = unsafe { &*ctx_ptr };
|
let ctx_ref = unsafe { &*ctx_ptr };
|
||||||
self.sampler.sample(ctx_ref, -1)
|
self.sampler.sample(ctx_ref, -1)
|
||||||
}
|
}
|
||||||
@@ -133,13 +136,39 @@ struct HealthResponse {
|
|||||||
model: String,
|
model: String,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// ── SSE Chunk Types ──
|
||||||
|
|
||||||
|
#[derive(Serialize)]
|
||||||
|
struct SseChunk {
|
||||||
|
id: String,
|
||||||
|
object: String,
|
||||||
|
created: i64,
|
||||||
|
model: String,
|
||||||
|
choices: Vec<SseChoice>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Serialize)]
|
||||||
|
struct SseChoice {
|
||||||
|
index: u32,
|
||||||
|
delta: SseDelta,
|
||||||
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
|
finish_reason: Option<String>,
|
||||||
|
}
|
||||||
|
|
||||||
|
#[derive(Serialize)]
|
||||||
|
struct SseDelta {
|
||||||
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
|
role: Option<String>,
|
||||||
|
#[serde(skip_serializing_if = "Option::is_none")]
|
||||||
|
content: Option<String>,
|
||||||
|
}
|
||||||
|
|
||||||
const DEFAULT_MODEL_PATH: &str = "/models/MiniCPM-V-4.6-Q4_K_M.gguf";
|
const DEFAULT_MODEL_PATH: &str = "/models/MiniCPM-V-4.6-Q4_K_M.gguf";
|
||||||
const EOS_TOKEN: i32 = 248044;
|
|
||||||
|
|
||||||
fn check_auth(headers: &HeaderMap) -> Result<(), (StatusCode, String)> {
|
fn check_auth(headers: &HeaderMap) -> Result<(), (StatusCode, String)> {
|
||||||
let api_key = std::env::var("API_KEY").unwrap_or_default();
|
let api_key = std::env::var("API_KEY").unwrap_or_default();
|
||||||
if api_key.is_empty() {
|
if api_key.is_empty() {
|
||||||
return Ok(()); // no key configured = open
|
return Ok(());
|
||||||
}
|
}
|
||||||
let header = headers
|
let header = headers
|
||||||
.get("authorization")
|
.get("authorization")
|
||||||
@@ -188,7 +217,6 @@ async fn main() {
|
|||||||
.new_context(&backend, ctx_params)
|
.new_context(&backend, ctx_params)
|
||||||
.expect("Failed to create context");
|
.expect("Failed to create context");
|
||||||
|
|
||||||
// Extend lifetime: model outlives context
|
|
||||||
let context: llama_cpp_2::context::LlamaContext<'static> =
|
let context: llama_cpp_2::context::LlamaContext<'static> =
|
||||||
unsafe { std::mem::transmute(context) };
|
unsafe { std::mem::transmute(context) };
|
||||||
|
|
||||||
@@ -238,13 +266,13 @@ async fn chat_completions(
|
|||||||
State(state): State<Arc<AppState>>,
|
State(state): State<Arc<AppState>>,
|
||||||
headers: HeaderMap,
|
headers: HeaderMap,
|
||||||
Json(req): Json<ChatRequest>,
|
Json(req): Json<ChatRequest>,
|
||||||
) -> Result<Json<ChatResponse>, (StatusCode, String)> {
|
) -> Result<impl IntoResponse, (StatusCode, String)> {
|
||||||
info!("Chat: {} chars, max_tokens={:?}", req.messages.len(), req.max_tokens);
|
|
||||||
check_auth(&headers)?;
|
check_auth(&headers)?;
|
||||||
let prompt = build_prompt(&req.messages);
|
|
||||||
let max_tokens = req.max_tokens.unwrap_or(256).min(1024);
|
|
||||||
|
|
||||||
info!("Chat: {} chars, max_tokens={}", prompt.len(), max_tokens);
|
let chat_id = format!("chatcmpl-{}", uuid::Uuid::new_v4());
|
||||||
|
let created = chrono::Utc::now().timestamp();
|
||||||
|
let max_tokens = req.max_tokens.unwrap_or(256).min(1024);
|
||||||
|
let prompt = build_prompt(&req.messages);
|
||||||
|
|
||||||
// Tokenize
|
// Tokenize
|
||||||
let input_tokens = state
|
let input_tokens = state
|
||||||
@@ -253,47 +281,138 @@ async fn chat_completions(
|
|||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
|
||||||
|
|
||||||
let prompt_tokens = input_tokens.len() as u32;
|
let prompt_tokens = input_tokens.len() as u32;
|
||||||
info!(" {} prompt tokens", prompt_tokens);
|
info!(" Chat: {} prompt tokens, max_tokens={}", prompt_tokens, max_tokens);
|
||||||
|
|
||||||
// Lock context
|
if req.stream.unwrap_or(false) {
|
||||||
|
// ── Streaming mode: spawn generator, pipe via mpsc channel ──
|
||||||
|
let state = state.clone();
|
||||||
|
let model_name = req.model.clone();
|
||||||
|
let (tx, rx) = tokio::sync::mpsc::channel::<Result<Event, std::convert::Infallible>>(64);
|
||||||
|
|
||||||
|
tokio::spawn(async move {
|
||||||
|
// Send role chunk
|
||||||
|
let role_chunk = serde_json::to_string(&SseChunk {
|
||||||
|
id: chat_id.clone(),
|
||||||
|
object: "chat.completion.chunk".into(),
|
||||||
|
created,
|
||||||
|
model: model_name.clone(),
|
||||||
|
choices: vec![SseChoice {
|
||||||
|
index: 0,
|
||||||
|
delta: SseDelta {
|
||||||
|
role: Some("assistant".into()),
|
||||||
|
content: None,
|
||||||
|
},
|
||||||
|
finish_reason: None,
|
||||||
|
}],
|
||||||
|
}).unwrap();
|
||||||
|
let _ = tx.send(Ok(Event::default().data(role_chunk))).await;
|
||||||
|
|
||||||
|
// Lock model context
|
||||||
let mut inner = state.ctx.lock().await;
|
let mut inner = state.ctx.lock().await;
|
||||||
|
inner.clear();
|
||||||
|
if let Err(e) = inner.prefill(&input_tokens) {
|
||||||
|
info!(" Prefill error: {e}");
|
||||||
|
return;
|
||||||
|
}
|
||||||
|
|
||||||
|
let mut count = 0u32;
|
||||||
|
loop {
|
||||||
|
if count >= max_tokens {
|
||||||
|
let chunk = serde_json::to_string(&SseChunk {
|
||||||
|
id: chat_id.clone(),
|
||||||
|
object: "chat.completion.chunk".into(),
|
||||||
|
created,
|
||||||
|
model: model_name.clone(),
|
||||||
|
choices: vec![SseChoice {
|
||||||
|
index: 0,
|
||||||
|
delta: SseDelta { role: None, content: None },
|
||||||
|
finish_reason: Some("length".into()),
|
||||||
|
}],
|
||||||
|
}).unwrap();
|
||||||
|
let _ = tx.send(Ok(Event::default().data(chunk))).await;
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
|
let token = inner.sample_token();
|
||||||
|
|
||||||
|
if state.model.is_eog_token(token) {
|
||||||
|
let chunk = serde_json::to_string(&SseChunk {
|
||||||
|
id: chat_id.clone(),
|
||||||
|
object: "chat.completion.chunk".into(),
|
||||||
|
created,
|
||||||
|
model: model_name.clone(),
|
||||||
|
choices: vec![SseChoice {
|
||||||
|
index: 0,
|
||||||
|
delta: SseDelta { role: None, content: None },
|
||||||
|
finish_reason: Some("stop".into()),
|
||||||
|
}],
|
||||||
|
}).unwrap();
|
||||||
|
let _ = tx.send(Ok(Event::default().data(chunk))).await;
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
|
||||||
|
let piece = decode_token_piece(&state.model, token);
|
||||||
|
let content = clean_text(&piece);
|
||||||
|
if !content.is_empty() {
|
||||||
|
let chunk = serde_json::to_string(&SseChunk {
|
||||||
|
id: chat_id.clone(),
|
||||||
|
object: "chat.completion.chunk".into(),
|
||||||
|
created,
|
||||||
|
model: model_name.clone(),
|
||||||
|
choices: vec![SseChoice {
|
||||||
|
index: 0,
|
||||||
|
delta: SseDelta { role: None, content: Some(content) },
|
||||||
|
finish_reason: None,
|
||||||
|
}],
|
||||||
|
}).unwrap();
|
||||||
|
if tx.send(Ok(Event::default().data(chunk))).await.is_err() {
|
||||||
|
break; // Client disconnected
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
let pos = input_tokens.len() as i32 + count as i32;
|
||||||
|
if let Err(e) = inner.decode_token(token, pos) {
|
||||||
|
info!(" Decode error: {e}");
|
||||||
|
break;
|
||||||
|
}
|
||||||
|
count += 1;
|
||||||
|
}
|
||||||
|
});
|
||||||
|
|
||||||
|
let stream = ReceiverStream::new(rx);
|
||||||
|
let sse = Sse::new(stream).keep_alive(KeepAlive::default());
|
||||||
|
Ok(sse.into_response())
|
||||||
|
} else {
|
||||||
|
// ── Non-streaming mode ──
|
||||||
|
let mut inner = state.ctx.lock().await;
|
||||||
inner.clear();
|
inner.clear();
|
||||||
inner.prefill(&input_tokens).map_err(|e| {
|
inner.prefill(&input_tokens).map_err(|e| {
|
||||||
(StatusCode::INTERNAL_SERVER_ERROR, format!("Prefill: {e}"))
|
(StatusCode::INTERNAL_SERVER_ERROR, format!("Prefill: {e}"))
|
||||||
})?;
|
})?;
|
||||||
|
|
||||||
// Generate
|
|
||||||
let mut output_tokens: Vec<LlamaToken> = Vec::new();
|
let mut output_tokens: Vec<LlamaToken> = Vec::new();
|
||||||
|
|
||||||
// First sample from prefill
|
|
||||||
let mut current = inner.sample_token();
|
let mut current = inner.sample_token();
|
||||||
|
|
||||||
for _ in 0..max_tokens {
|
for _ in 0..max_tokens {
|
||||||
if current.0 == EOS_TOKEN {
|
if state.model.is_eog_token(current) {
|
||||||
break;
|
break;
|
||||||
}
|
}
|
||||||
let pos = input_tokens.len() as i32 + output_tokens.len() as i32;
|
let pos = input_tokens.len() as i32 + output_tokens.len() as i32;
|
||||||
output_tokens.push(current);
|
output_tokens.push(current);
|
||||||
|
|
||||||
// Decode the last token
|
|
||||||
inner.decode_token(current, pos)
|
inner.decode_token(current, pos)
|
||||||
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("Decode: {e}")))?;
|
.map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, format!("Decode: {e}")))?;
|
||||||
|
|
||||||
// Sample next token
|
|
||||||
current = inner.sample_token();
|
current = inner.sample_token();
|
||||||
}
|
}
|
||||||
|
|
||||||
let output_text = decode_tokens(&state.model, &output_tokens);
|
let output_text = decode_tokens(&state.model, &output_tokens);
|
||||||
|
|
||||||
let completion_tokens = output_tokens.len() as u32;
|
let completion_tokens = output_tokens.len() as u32;
|
||||||
|
|
||||||
info!(" {} generated tokens", completion_tokens);
|
info!(" {} generated tokens", completion_tokens);
|
||||||
|
|
||||||
Ok(Json(ChatResponse {
|
Ok(Json(ChatResponse {
|
||||||
id: format!("chatcmpl-{}", uuid::Uuid::new_v4()),
|
id: chat_id,
|
||||||
object: "chat.completion".into(),
|
object: "chat.completion".into(),
|
||||||
created: chrono::Utc::now().timestamp(),
|
created,
|
||||||
model: req.model,
|
model: req.model,
|
||||||
choices: vec![Choice {
|
choices: vec![Choice {
|
||||||
index: 0,
|
index: 0,
|
||||||
@@ -301,51 +420,64 @@ async fn chat_completions(
|
|||||||
role: "assistant".into(),
|
role: "assistant".into(),
|
||||||
content: output_text,
|
content: output_text,
|
||||||
},
|
},
|
||||||
finish_reason: if completion_tokens < max_tokens {
|
finish_reason: if completion_tokens < max_tokens { "stop" } else { "length" }.into(),
|
||||||
"stop"
|
|
||||||
} else {
|
|
||||||
"length"
|
|
||||||
}
|
|
||||||
.into(),
|
|
||||||
}],
|
}],
|
||||||
usage: Usage {
|
usage: Usage {
|
||||||
prompt_tokens,
|
prompt_tokens,
|
||||||
completion_tokens,
|
completion_tokens,
|
||||||
total_tokens: prompt_tokens + completion_tokens,
|
total_tokens: prompt_tokens + completion_tokens,
|
||||||
},
|
},
|
||||||
}))
|
})
|
||||||
|
.into_response())
|
||||||
}
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// ── Helper Functions ──
|
||||||
|
|
||||||
fn build_prompt(messages: &[ChatMessage]) -> String {
|
fn build_prompt(messages: &[ChatMessage]) -> String {
|
||||||
let mut prompt = String::new();
|
let mut prompt = String::new();
|
||||||
for msg in messages {
|
for (i, msg) in messages.iter().enumerate() {
|
||||||
match msg.role.as_str() {
|
let role = match msg.role.as_str() {
|
||||||
"system" => prompt.push_str(&format!("System: {}\n", msg.content)),
|
"system" => "system",
|
||||||
"user" => prompt.push_str(&format!("User: {}\n", msg.content)),
|
"user" => "user",
|
||||||
"assistant" => prompt.push_str(&format!("Assistant: {}\n", msg.content)),
|
"assistant" => "assistant",
|
||||||
_ => prompt.push_str(&format!("{}\n", msg.content)),
|
_ => "user",
|
||||||
|
};
|
||||||
|
if i == 0 && role == "system" {
|
||||||
|
prompt.push_str(&format!("<|im_start|>system\n{}<|im_end|>\n", msg.content));
|
||||||
|
} else {
|
||||||
|
prompt.push_str(&format!("<|im_start|>{}\n{}<|im_end|>\n", role, msg.content));
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
prompt.push_str("Assistant: ");
|
prompt.push_str("<|im_start|>assistant\n<think>\n\n</think>\n\n");
|
||||||
prompt
|
prompt
|
||||||
}
|
}
|
||||||
|
|
||||||
|
fn decode_token_piece(model: &LlamaModel, token: LlamaToken) -> String {
|
||||||
|
let bytes = match model.token_to_piece_bytes(token, 32, true, None) {
|
||||||
|
Ok(b) => b,
|
||||||
|
Err(TokenToStringError::InsufficientBufferSpace(neg)) => {
|
||||||
|
let size = (-neg).max(0).try_into().unwrap_or(256);
|
||||||
|
model.token_to_piece_bytes(token, size, true, None).unwrap_or_default()
|
||||||
|
}
|
||||||
|
_ => return String::new(),
|
||||||
|
};
|
||||||
|
String::from_utf8(bytes).unwrap_or_default()
|
||||||
|
}
|
||||||
|
|
||||||
|
fn clean_text(text: &str) -> String {
|
||||||
|
text.replace("<|im_end|>", "")
|
||||||
|
.replace("<|im_start|>", "")
|
||||||
|
.replace("<think>", "")
|
||||||
|
.replace("</think>", "")
|
||||||
|
.trim()
|
||||||
|
.to_string()
|
||||||
|
}
|
||||||
|
|
||||||
fn decode_tokens(model: &LlamaModel, tokens: &[LlamaToken]) -> String {
|
fn decode_tokens(model: &LlamaModel, tokens: &[LlamaToken]) -> String {
|
||||||
let mut out = String::with_capacity(tokens.len() * 4);
|
let mut out = String::with_capacity(tokens.len() * 4);
|
||||||
for &token in tokens {
|
for &token in tokens {
|
||||||
// Try with a reasonable initial buffer (32 bytes)
|
out.push_str(&decode_token_piece(model, token));
|
||||||
let bytes = match model.token_to_piece_bytes(token, 32, true, None) {
|
|
||||||
Ok(b) => b,
|
|
||||||
Err(TokenToStringError::InsufficientBufferSpace(neg)) => {
|
|
||||||
// Retry with the suggested buffer size
|
|
||||||
let size = (-neg).max(0).try_into().unwrap_or(256);
|
|
||||||
model.token_to_piece_bytes(token, size, true, None).unwrap_or_default()
|
|
||||||
}
|
}
|
||||||
_ => continue,
|
clean_text(&out)
|
||||||
};
|
|
||||||
if let Ok(s) = String::from_utf8(bytes) {
|
|
||||||
out.push_str(&s);
|
|
||||||
}
|
|
||||||
}
|
|
||||||
out
|
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user