//! QQQ (Quick Query Quorum) is deliberately a single-file, in-memory group-chat server.
//!
//! The global map lock is held only long enough to locate a session; code then acquires that
//! session's lock and never performs provider I/O while holding it. History is append-only: an
//! assistant tool call and its matching tool result are admitted as one pair. Pending inputs are
//! drained in arrival order into one serial model loop, while deliveries remain FIFO with one
//! active item at a time. Participant tokens are credentials, never public view data.
// Configuration, errors, and transport DTOs.
use axum::{
Json, Router,
body::{Body, to_bytes},
extract::{
DefaultBodyLimit, FromRequest, Multipart, Path, Request, State,
ws::{Message, WebSocket, WebSocketUpgrade},
},
http::{HeaderMap, StatusCode, header},
response::{Html, IntoResponse, Response},
routing::{get, post},
};
use base64::{Engine as _, engine::general_purpose::STANDARD};
use futures_util::StreamExt;
use reqwest::header::{ACCEPT, AUTHORIZATION, CONTENT_LENGTH, CONTENT_TYPE};
use serde::{Deserialize, Serialize};
use serde_json::{Value, json};
use std::{
collections::{HashMap, HashSet, VecDeque},
env,
sync::Arc,
time::{Duration, Instant},
};
use tokio::sync::{Mutex, mpsc};
use uuid::Uuid;
const DEFAULT_PRIVATE_CONVERSATION_PROMPT: &str = "You are a helpful private assistant. Continue a normal one-to-one conversation. Create a group only when the user explicitly requests one and gives a concrete purpose. A display name and language are optional. When creating it, generate an expanded, use-case-specific session_prompt from that purpose. Never infer group intent from an ordinary request.";
const DEFAULT_GROUP_SESSION_PROTOCOL_PROMPT: &str = "Transport and protocol rules: participants have stable API names such as user_1 and user_2; a display name, if present, is repeated in their message. Respond only through send_messages, addressing current API names or \"all\". Keep deliveries concise. Each dispatch independently has next_speaker; use it only when this session calls for a floor change after that delivery, otherwise use null.";
const MAX_JSON_BYTES: usize = 16 * 1024;
/// Enough for boundaries and disposition headers while keeping multipart requests bounded.
const MAX_MULTIPART_OVERHEAD: usize = 16 * 1024;
#[derive(Clone)]
struct Config {
bind: String,
base_url: String,
api_key: String,
chat_model: String,
stt_model: Option<String>,
tts: Option<TtsConfig>,
referer: Option<String>,
title: Option<String>,
max_sessions: usize,
max_participants: usize,
ttl: Duration,
max_history: usize,
max_audio: usize,
batch: Duration,
delivery_timeout: Duration,
default_language: String,
private_conversation_prompt: String,
group_session_protocol_prompt: String,
}
#[derive(Clone)]
struct TtsConfig {
tts_model: String,
tts_voice: String,
tts_format: String,
}
impl Config {
fn load() -> Result<Self, String> {
let required = |k: &str| env::var(k).map_err(|_| format!("missing {k}"));
let value = |k: &str, d: &str| env::var(k).unwrap_or_else(|_| d.into());
let number = |k: &str, d: usize| parse_positive(&value(k, &d.to_string()), k);
let optional = |k: &str| env::var(k).unwrap_or_default().trim().to_string();
let stt_model = optional("STT_MODEL");
let tts_model = optional("TTS_MODEL");
let tts_voice = optional("TTS_VOICE");
let (stt_model, tts) =
voice_config(stt_model, tts_model, tts_voice, value("TTS_FORMAT", "mp3"))?;
Ok(Self {
bind: required_trimmed("QQQ_BIND", &value("QQQ_BIND", "127.0.0.1:8787"))?,
base_url: required_trimmed(
"OPENAI_BASE_URL",
&value("OPENAI_BASE_URL", "https://openrouter.ai/api/v1"),
)?
.trim_end_matches('/')
.into(),
api_key: required_trimmed("OPENAI_API_KEY", &required("OPENAI_API_KEY")?)?,
chat_model: required_trimmed("CHAT_MODEL", &required("CHAT_MODEL")?)?,
stt_model,
tts,
referer: env::var("OPENAI_HTTP_REFERER")
.ok()
.filter(|s| !s.is_empty()),
title: env::var("OPENAI_APP_TITLE").ok().filter(|s| !s.is_empty()),
max_sessions: number("MAX_SESSIONS", 100)?,
max_participants: number("MAX_PARTICIPANTS_PER_SESSION", 8)?,
ttl: Duration::from_secs(number("SESSION_TTL_SECONDS", 14400)? as u64),
max_history: number("MAX_HISTORY_MESSAGES", 512)?,
max_audio: number("MAX_AUDIO_BYTES", 10_000_000)?,
batch: Duration::from_millis(number("INPUT_BATCH_MS", 350)? as u64),
delivery_timeout: Duration::from_secs(number("DELIVERY_TIMEOUT_SECONDS", 45)? as u64),
default_language: required_trimmed(
"DEFAULT_LANGUAGE",
&value("DEFAULT_LANGUAGE", "en"),
)?,
private_conversation_prompt: value(
"PRIVATE_CONVERSATION_PROMPT",
DEFAULT_PRIVATE_CONVERSATION_PROMPT,
),
group_session_protocol_prompt: value(
"GROUP_SESSION_PROTOCOL_PROMPT",
DEFAULT_GROUP_SESSION_PROTOCOL_PROMPT,
),
})
}
}
fn required_trimmed(name: &str, value: &str) -> Result<String, String> {
let value = value.trim();
(!value.is_empty())
.then(|| value.to_string())
.ok_or_else(|| format!("missing {name}"))
}
fn parse_positive(value: &str, name: &str) -> Result<usize, String> {
value
.parse::<usize>()
.map_err(|_| format!("invalid {name}"))?
.checked_sub(1)
.map(|value| value + 1)
.ok_or_else(|| format!("invalid {name}"))
}
fn voice_config(
stt_model: String,
tts_model: String,
tts_voice: String,
tts_format: String,
) -> Result<(Option<String>, Option<TtsConfig>), String> {
let tts = match (tts_model.is_empty(), tts_voice.is_empty()) {
(true, true) => None,
(false, false) => Some(TtsConfig {
tts_model,
tts_voice,
tts_format: supported_tts_format(&tts_format)?,
}),
_ => return Err("TTS_MODEL and TTS_VOICE must either both be set or both be empty".into()),
};
Ok(((!stt_model.is_empty()).then_some(stt_model), tts))
}
fn supported_tts_format(value: &str) -> Result<String, String> {
let value = required_trimmed("TTS_FORMAT", value)?;
matches!(value.as_str(), "mp3" | "wav" | "ogg" | "webm" | "m4a")
.then_some(value)
.ok_or_else(|| "invalid TTS_FORMAT".into())
}
fn startup_config_error(error: &str) -> String {
format!("configuration error: {error}; set it in .env or export it")
}
#[derive(Debug)]
struct ApiError(&'static str, &'static str, StatusCode);
impl IntoResponse for ApiError {
fn into_response(self) -> Response {
(
self.2,
Json(json!({"error":{"code":self.0,"message":self.1}})),
)
.into_response()
}
}
type ApiResult<T> = Result<T, ApiError>;
async fn json_body<T: serde::de::DeserializeOwned>(request: Request) -> ApiResult<T> {
if !request
.headers()
.get(header::CONTENT_TYPE)
.and_then(|value| value.to_str().ok())
.is_some_and(|value| value.split(';').next() == Some("application/json"))
{
return Err(bad());
}
let bytes = to_bytes(request.into_body(), MAX_JSON_BYTES)
.await
.map_err(|_| bad())?;
serde_json::from_slice(&bytes).map_err(|_| bad())
}
fn bad() -> ApiError {
ApiError(
"invalid_request",
"The request could not be completed.",
StatusCode::BAD_REQUEST,
)
}
fn limited_text(text: String) -> ApiResult<String> {
let text = text.trim().to_string();
if text.is_empty() || text.chars().count() > 4000 {
Err(bad())
} else {
Ok(text)
}
}
fn provider() -> ApiError {
ApiError(
"provider_failure",
"The request could not be completed.",
StatusCode::BAD_GATEWAY,
)
}
fn unauthorized() -> ApiError {
ApiError(
"unauthorized",
"The request could not be completed.",
StatusCode::UNAUTHORIZED,
)
}
fn session_not_found() -> ApiError {
ApiError(
"session_not_found",
"The request could not be completed.",
StatusCode::NOT_FOUND,
)
}
fn chat_failure() -> ApiError {
eprintln!("chat failed");
provider()
}
// provider/history DTOs
#[derive(Clone, Serialize, Deserialize, Debug)]
struct ChatMessage {
role: String,
#[serde(skip_serializing_if = "Option::is_none")]
name: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
content: Option<ChatContent>,
#[serde(skip_serializing_if = "Option::is_none")]
tool_calls: Option<Vec<ToolCall>>,
#[serde(skip_serializing_if = "Option::is_none")]
tool_call_id: Option<String>,
}
#[derive(Clone, Serialize, Deserialize, Debug, PartialEq)]
#[serde(untagged)]
enum ChatContent {
Text(String),
Parts(Vec<ContentPart>),
}
#[derive(Clone, Serialize, Deserialize, Debug, PartialEq)]
#[serde(tag = "type")]
enum ContentPart {
#[serde(rename = "text")]
Text { text: String },
#[serde(rename = "input_audio")]
InputAudio { input_audio: AudioPart },
}
#[derive(Clone, Serialize, Deserialize, Debug, PartialEq)]
struct AudioPart {
data: String,
format: AudioFormat,
}
#[derive(Clone, Serialize, Deserialize, Debug, PartialEq)]
#[serde(rename_all = "lowercase")]
enum AudioFormat {
Webm,
Ogg,
M4a,
Mp3,
Wav,
}
impl AudioFormat {
fn from_mime(mime: &str) -> Option<Self> {
match mime.split(';').next().unwrap_or("").trim() {
"audio/webm" => Some(Self::Webm),
"audio/ogg" => Some(Self::Ogg),
"audio/mp4" => Some(Self::M4a),
"audio/mpeg" => Some(Self::Mp3),
"audio/wav" => Some(Self::Wav),
_ => None,
}
}
fn as_str(&self) -> &'static str {
match self {
Self::Webm => "webm",
Self::Ogg => "ogg",
Self::M4a => "m4a",
Self::Mp3 => "mp3",
Self::Wav => "wav",
}
}
}
impl ChatContent {
#[cfg(test)]
fn as_text(&self) -> Option<&str> {
match self {
Self::Text(text) => Some(text),
Self::Parts(_) => None,
}
}
fn text(self) -> Option<String> {
match self {
Self::Text(text) => Some(text),
Self::Parts(_) => None,
}
}
}
impl From<String> for ChatContent {
fn from(value: String) -> Self {
Self::Text(value)
}
}
impl From<&str> for ChatContent {
fn from(value: &str) -> Self {
Self::Text(value.into())
}
}
#[derive(Clone, Serialize, Deserialize, Debug)]
struct ToolCall {
id: String,
#[serde(rename = "type")]
kind: String,
function: ToolFunction,
}
#[derive(Clone, Serialize, Deserialize, Debug)]
struct ToolFunction {
name: String,
arguments: String,
}
#[derive(Deserialize, Debug)]
struct ProviderReply {
choices: Vec<Choice>,
}
#[derive(Deserialize, Debug)]
struct Choice {
finish_reason: Option<String>,
message: ProviderMessage,
}
#[derive(Deserialize, Debug)]
struct ProviderMessage {
role: String,
tool_calls: Option<Vec<ToolCall>>,
content: Option<ChatContent>,
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct SendMessagesArgs {
dispatches: Vec<Dispatch>,
}
#[derive(Deserialize, Clone)]
#[serde(deny_unknown_fields)]
struct Dispatch {
reply_to: String,
text: String,
next_speaker: Option<String>,
}
#[derive(Serialize, Deserialize)]
struct TtsRequest {
model: String,
voice: String,
input: String,
response_format: String,
}
// domain state
struct AppState {
config: Config,
http: reqwest::Client,
sessions: Mutex<HashMap<String, Arc<Mutex<Session>>>>,
}
struct Session {
id: String,
last_activity_at: Instant,
language: String,
status: SessionStatus,
participants: HashMap<String, Participant>,
participant_order: Vec<String>,
messages: Vec<ChatMessage>,
onboarding_messages: Vec<ChatMessage>,
pending_messages: VecDeque<ChatMessage>,
model_running: bool,
debounce_scheduled: bool,
deliveries: VecDeque<Delivery>,
active_delivery: Option<Delivery>,
floor_holder: Option<String>,
sockets: HashMap<String, SocketConnection>,
onboarding_audio: Option<OnboardingAudio>,
onboarding_revision: u64,
input_revision: u64,
delivery_acks: HashMap<String, HashSet<String>>,
}
#[derive(Debug, PartialEq)]
enum SessionStatus {
Onboarding,
Active,
Closed,
}
struct Participant {
id: String,
token: String,
api_name: String,
display_name: Option<String>,
browser_language: String,
onboarding: Onboarding,
connected: bool,
creator: bool,
}
#[derive(PartialEq)]
enum Onboarding {
Required,
Complete,
}
struct Delivery {
id: String,
participant_id: String,
participant_api_name: String,
text: String,
audio: Option<Vec<u8>>,
mime: String,
next_speaker: Option<String>,
ready: bool,
}
struct OnboardingAudio {
id: String,
bytes: Option<Vec<u8>>,
mime: String,
text: String,
}
struct SocketConnection {
id: String,
tx: mpsc::Sender<ServerEvent>,
}
#[derive(Default)]
struct EmitOutcome {
start_delivery: bool,
}
impl EmitOutcome {
fn merge(&mut self, other: Self) {
self.start_delivery |= other.start_delivery;
}
}
#[derive(Clone, Serialize)]
struct ServerEvent {
#[serde(rename = "type")]
kind: String,
#[serde(flatten)]
data: Value,
}
impl ServerEvent {
fn new(kind: &str, data: Value) -> Self {
Self {
kind: kind.into(),
data,
}
}
}
/// Send an event without ever waiting on a slow WebSocket writer.
fn send_event(s: &mut Session, participant: Option<&str>, event: ServerEvent) {
let mut saturated = Vec::new();
for (id, connection) in &s.sockets {
if participant.map(|x| x == id).unwrap_or(true)
&& connection.tx.try_send(event.clone()).is_err()
{
saturated.push(id.clone());
}
}
for id in saturated {
s.sockets.remove(&id);
if let Some(participant) = s.participants.get_mut(&id) {
participant.connected = false;
}
for acknowledgements in s.delivery_acks.values_mut() {
acknowledgements.remove(&id);
}
}
}
/// Complete an active delivery whose recipients all became unavailable. This transition is
/// synchronous and idempotent; callers only need to trigger its next FIFO item after unlocking.
fn complete_unacknowledgeable_delivery(s: &mut Session) -> bool {
let Some(delivery_id) = s.active_delivery.as_ref().and_then(|active| {
s.delivery_acks
.get(&active.id)
.is_some_and(HashSet::is_empty)
.then(|| active.id.clone())
}) else {
return false;
};
let active = s.active_delivery.take().expect("active delivery checked");
s.delivery_acks.remove(&delivery_id);
s.last_activity_at = Instant::now();
send_event(
s,
None,
ServerEvent::new("delivery.completed", json!({"delivery_id":delivery_id})),
);
if let Some(next) = active.next_speaker {
s.floor_holder = s
.participants
.values()
.find(|p| p.api_name == next && s.sockets.contains_key(&p.id))
.map(|p| p.id.clone());
let floor_holder = s
.floor_holder
.as_ref()
.and_then(|id| s.participants.get(id))
.map(|p| p.api_name.clone());
send_event(
s,
None,
ServerEvent::new("floor.changed", json!({"floor_holder":floor_holder})),
);
}
!s.deliveries.is_empty()
}
/// Saturation removal and active-delivery completion are one state transition.
fn emit(s: &mut Session, participant: Option<&str>, event: ServerEvent) -> EmitOutcome {
send_event(s, participant, event);
EmitOutcome {
start_delivery: complete_unacknowledgeable_delivery(s),
}
}
fn set_floor(s: &mut Session, participant_id: &str, required: bool) {
let changed = s.floor_holder.as_deref() != Some(participant_id);
s.floor_holder = Some(participant_id.to_string());
if required || changed {
let api_name = s
.participants
.get(participant_id)
.map(|p| p.api_name.clone());
emit(
s,
None,
ServerEvent::new("floor.changed", json!({"floor_holder":api_name})),
);
}
}
fn interrupt_deliveries(s: &mut Session, participant_id: &str) {
let had_deliveries = s.active_delivery.is_some() || !s.deliveries.is_empty();
if let Some(active) = s.active_delivery.take() {
s.delivery_acks.remove(&active.id);
emit(
s,
None,
ServerEvent::new("delivery.interrupted", json!({"delivery_id":active.id})),
);
}
s.deliveries.clear();
set_floor(s, participant_id, had_deliveries);
}
fn session_expired(s: &Session, ttl: Duration) -> bool {
s.sockets.is_empty() && s.last_activity_at.elapsed() > ttl
}
fn owns_socket(s: &Session, participant_id: &str, connection_id: &str) -> bool {
s.sockets
.get(participant_id)
.is_some_and(|connection| connection.id == connection_id)
}
// validation/state transition helpers
fn now_session(id: String) -> Session {
Session {
id,
last_activity_at: Instant::now(),
language: String::new(),
status: SessionStatus::Onboarding,
participants: HashMap::new(),
participant_order: vec![],
messages: vec![],
onboarding_messages: vec![],
pending_messages: VecDeque::new(),
model_running: false,
debounce_scheduled: false,
deliveries: VecDeque::new(),
active_delivery: None,
floor_holder: None,
sockets: HashMap::new(),
onboarding_audio: None,
onboarding_revision: 1,
input_revision: 0,
delivery_acks: HashMap::new(),
}
}
fn api_name(s: &Session) -> String {
format!("user_{}", s.participant_order.len() + 1)
}
fn add_participant(s: &mut Session, language: String, creator: bool) -> Participant {
let p = Participant {
id: Uuid::new_v4().to_string(),
token: Uuid::new_v4().to_string(),
api_name: api_name(s),
display_name: None,
browser_language: language,
connected: false,
onboarding: if creator {
Onboarding::Required
} else {
Onboarding::Complete
},
creator,
};
s.participant_order.push(p.id.clone());
let returned = p.clone_for_return();
s.participants.insert(p.id.clone(), p);
returned
}
impl Participant {
fn clone_for_return(&self) -> Self {
Self {
id: self.id.clone(),
token: self.token.clone(),
api_name: self.api_name.clone(),
display_name: self.display_name.clone(),
browser_language: self.browser_language.clone(),
onboarding: if self.onboarding == Onboarding::Required {
Onboarding::Required
} else {
Onboarding::Complete
},
connected: self.connected,
creator: self.creator,
}
}
}
fn participant_by_token<'a>(s: &'a Session, token: &str) -> Option<&'a Participant> {
s.participants.values().find(|p| p.token == token)
}
fn require_auth<'a>(s: &'a Session, token: &str) -> ApiResult<&'a Participant> {
participant_by_token(s, token).ok_or_else(unauthorized)
}
async fn lookup_session(st: &AppState, id: &str) -> ApiResult<Arc<Mutex<Session>>> {
st.sessions
.lock()
.await
.get(id)
.cloned()
.ok_or_else(session_not_found)
}
fn history_room(s: &mut Session, max: usize, count: usize) -> ApiResult<()> {
if s.messages.len() + count > max {
s.status = SessionStatus::Closed;
emit(
s,
None,
ServerEvent::new("error", json!({"code":"history_limit"})),
);
return Err(ApiError(
"history_limit",
"The request could not be completed.",
StatusCode::CONFLICT,
));
}
Ok(())
}
fn user_message(p: &Participant, text: String) -> ChatMessage {
ChatMessage {
role: "user".into(),
name: Some(p.api_name.clone()),
content: Some(ChatContent::Text(format!(
"{}: {}",
p.display_name.as_deref().unwrap_or(&p.api_name),
text.trim()
))),
tool_calls: None,
tool_call_id: None,
}
}
fn event_message(text: String) -> ChatMessage {
ChatMessage {
role: "user".into(),
name: Some("qqq_event".into()),
content: Some(ChatContent::Text(text)),
tool_calls: None,
tool_call_id: None,
}
}
/// Reads exactly one supported browser audio field; the route also bounds multipart framing.
async fn browser_audio(request: Request, limit: usize) -> ApiResult<(Vec<u8>, AudioFormat)> {
let total_limit = limit.saturating_add(MAX_MULTIPART_OVERHEAD);
let (parts, body) = request.into_parts();
let body = to_bytes(body, total_limit).await.map_err(|_| {
ApiError(
"audio_too_large",
"The request could not be completed.",
StatusCode::PAYLOAD_TOO_LARGE,
)
})?;
let request = Request::from_parts(parts, Body::from(body));
let mut multipart = Multipart::from_request(request, &())
.await
.map_err(|_| bad())?;
let mut audio = None;
while let Some(mut field) = multipart.next_field().await.map_err(|_| bad())? {
if field.name() != Some("audio") || audio.is_some() {
return Err(bad());
}
let format = field
.content_type()
.and_then(AudioFormat::from_mime)
.ok_or_else(bad)?;
let mut bytes = Vec::new();
while let Some(chunk) = field.next().await {
let chunk = chunk.map_err(|_| bad())?;
if bytes.len().saturating_add(chunk.len()) > limit {
return Err(ApiError(
"audio_too_large",
"The request could not be completed.",
StatusCode::PAYLOAD_TOO_LARGE,
));
}
bytes.extend_from_slice(&chunk);
}
if bytes.is_empty() {
return Err(bad());
}
audio = Some((bytes, format));
}
audio.ok_or_else(bad)
}
fn validate_tool(s: &Session, args: SendMessagesArgs) -> Result<Vec<Dispatch>, ()> {
if args.dispatches.is_empty() || args.dispatches.len() > 16 {
return Err(());
}
let mut out = vec![];
for d in args.dispatches {
let text = d.text.trim().to_string();
if text.is_empty() || text.chars().count() > 2000 {
return Err(());
}
if d.reply_to == "all" {
if !s
.participants
.values()
.any(|p| p.onboarding == Onboarding::Complete && s.sockets.contains_key(&p.id))
{
return Err(());
}
out.push(Dispatch {
reply_to: d.reply_to,
text,
next_speaker: d.next_speaker,
});
} else if s.participants.values().any(|p| {
p.onboarding == Onboarding::Complete
&& p.api_name == d.reply_to
&& s.sockets.contains_key(&p.id)
}) {
out.push(Dispatch {
reply_to: d.reply_to,
text,
next_speaker: d.next_speaker,
});
} else {
return Err(());
}
}
if out.is_empty() {
return Err(());
}
if out.iter().filter_map(|d| d.next_speaker.as_ref()).any(|n| {
!s.participants.values().any(|p| {
p.onboarding == Onboarding::Complete
&& p.api_name == *n
&& s.sockets.contains_key(&p.id)
})
}) {
return Err(());
}
Ok(out)
}
fn parse_tool(reply: ProviderReply) -> Result<(ToolCall, SendMessagesArgs), ()> {
let choice = reply.choices.into_iter().next().ok_or(())?;
if choice.finish_reason.as_deref() != Some("tool_calls")
|| choice.message.role != "assistant"
|| choice.message.content.is_some()
{
return Err(());
}
let mut calls = choice.message.tool_calls.ok_or(())?;
if calls.len() != 1 {
return Err(());
}
let call = calls.pop().ok_or(())?;
if call.kind != "function" || call.function.name != "send_messages" || call.id.is_empty() {
return Err(());
}
let args = serde_json::from_str(&call.function.arguments).map_err(|_| ())?;
Ok((call, args))
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct StartGroupSessionArgs {
display_name: Option<String>,
language: Option<String>,
session_prompt: String,
}
fn language_template(language: &str) -> &'static str {
if language.to_lowercase().starts_with("ru") {
"О чём вы хотите поговорить?"
} else {
"What would you like to do?"
}
}
fn group_system_prompt(session_prompt: &str, protocol_prompt: &str) -> String {
format!("{session_prompt}\n\n{protocol_prompt}")
}
fn private_system_prompt(config: &Config) -> String {
config.private_conversation_prompt.clone()
}
struct DisconnectOutcome {
active_delivery: Option<String>,
start_delivery: bool,
queued_event: bool,
}
fn active_delivery_target_connected(s: &Session, delivery_id: &str) -> bool {
s.active_delivery
.as_ref()
.filter(|d| d.id == delivery_id)
.and_then(|d| s.delivery_acks.get(&d.id))
.is_some_and(|ids| ids.iter().any(|id| s.sockets.contains_key(id)))
}
fn next_delivery_ready(s: &Session) -> bool {
s.deliveries.front().is_some_and(|delivery| delivery.ready)
}
fn model_result_is_current(s: &Session, revision: u64) -> bool {
s.input_revision == revision
}
/// The public view is intentionally assembled from authority identifiers only; credentials,
/// history, provider payloads, and queued audio never cross this boundary.
fn public_view(s: &Session) -> Value {
let participants = s
.participant_order
.iter()
.filter_map(|id| s.participants.get(id))
.map(|p| {
json!({"api_name":p.api_name,"display_name":p.display_name,"connected":p.connected,"has_floor":s.floor_holder.as_deref()==Some(&p.id)})
})
.collect::<Vec<_>>();
json!({"session_id":s.id,"status":if s.status==SessionStatus::Active{"active"}else if s.status==SessionStatus::Closed{"closed"}else{"onboarding"},"language":s.language,"participants":participants,"floor_holder":s.floor_holder.as_ref().and_then(|id|s.participants.get(id)).map(|p|p.api_name.clone()),"ai_state":if s.model_running{"thinking"}else{"idle"},"active_delivery":s.active_delivery.as_ref().map(|d|json!({"target":d.participant_api_name}))})
}
fn disconnect_participant(
s: &mut Session,
participant_id: &str,
max_history: usize,
) -> DisconnectOutcome {
let Some(p) = s.participants.get(participant_id) else {
return DisconnectOutcome {
active_delivery: None,
start_delivery: false,
queued_event: false,
};
};
let was_connected = p.connected;
let already_disconnected = !was_connected && !s.sockets.contains_key(participant_id);
let mut active_delivery = s
.active_delivery
.as_ref()
.filter(|d| d.participant_id == participant_id)
.map(|d| d.id.clone());
if already_disconnected {
return DisconnectOutcome {
active_delivery,
start_delivery: false,
queued_event: false,
};
}
let api = p.api_name.clone();
let name = p.display_name.clone();
if let Some(p) = s.participants.get_mut(participant_id) {
p.connected = false;
}
s.sockets.remove(participant_id);
if let Some(active) = s.active_delivery.as_ref()
&& let Some(acks) = s.delivery_acks.get_mut(&active.id)
{
acks.remove(participant_id);
}
if let Some(active) = s.active_delivery.as_ref()
&& s.delivery_acks
.get(&active.id)
.is_some_and(HashSet::is_empty)
{
active_delivery = Some(active.id.clone());
}
if s.floor_holder.as_deref() == Some(participant_id) {
s.floor_holder = None;
}
let queued_event =
was_connected && s.status == SessionStatus::Active && s.messages.len() < max_history;
if queued_event {
s.pending_messages.push_back(event_message(format!(
"Participant left: {}{}.",
api,
name.map_or_else(String::new, |name| format!(" ({name})"))
)));
}
let mut start_delivery = false;
if was_connected {
let transition = emit(
s,
None,
ServerEvent::new("participant.left", json!({"api_name":api})),
);
if transition.start_delivery {
active_delivery = None;
start_delivery = true;
}
}
s.last_activity_at = Instant::now();
DisconnectOutcome {
active_delivery,
start_delivery,
queued_event,
}
}
// provider functions chat/transcribe/speak
fn request_headers(c: &Config, request: reqwest::RequestBuilder) -> reqwest::RequestBuilder {
let mut r = request.header(AUTHORIZATION, format!("Bearer {}", c.api_key));
if let Some(v) = &c.referer {
r = r.header("HTTP-Referer", v);
}
if let Some(v) = &c.title {
r = r.header("X-OpenRouter-Title", v);
}
r
}
fn tts_request(
client: &reqwest::Client,
config: &Config,
input: &str,
) -> Result<reqwest::RequestBuilder, serde_json::Error> {
let voice = config
.tts
.as_ref()
.expect("voice configuration checked by caller");
let payload = serde_json::to_vec(&TtsRequest {
model: voice.tts_model.clone(),
voice: voice.tts_voice.clone(),
input: input.into(),
response_format: voice.tts_format.clone(),
})?;
let content_length = payload.len().to_string();
Ok(request_headers(
config,
client
.post(format!("{}/audio/speech", config.base_url))
.header(ACCEPT, "audio/*")
.header(CONTENT_TYPE, "application/json")
.header(CONTENT_LENGTH, content_length)
.body(payload),
))
}
fn tool_schema() -> Value {
json!({"type":"function","function":{"name":"send_messages","description":"Queues concise AI messages for current participants. Each dispatch may independently change the speaking floor when that delivery completes; use null when no floor change is appropriate.","parameters":{"type":"object","additionalProperties":false,"properties":{"dispatches":{"type":"array","minItems":1,"maxItems":16,"items":{"type":"object","additionalProperties":false,"properties":{"reply_to":{"type":"string","description":"A current stable participant API name such as user_1, or the literal all."},"text":{"type":"string","minLength":1,"maxLength":2000,"description":"The concise message to deliver."},"next_speaker":{"type":["string","null"],"description":"A current participant API name to receive the floor as soon as this delivery completes, or null."}},"required":["reply_to","text","next_speaker"]}}},"required":["dispatches"]}}})
}
async fn chat(
state: &AppState,
messages: Vec<ChatMessage>,
tool: Value,
force_tool: bool,
session_id: &str,
) -> ApiResult<ProviderReply> {
eprintln!("chat started messages={}", messages.len());
let tool_choice = if force_tool {
json!({"type":"function","function":{"name":tool["function"]["name"]}})
} else {
json!("auto")
};
let messages = serde_json::to_value(messages).map_err(|_| provider())?;
let body = json!({"model":state.config.chat_model,"session_id":session_id,"messages":messages,"tools":[tool],"tool_choice":tool_choice,"parallel_tool_calls":false,"stream":false});
let res = match request_headers(
&state.config,
state
.http
.post(format!("{}/chat/completions", state.config.base_url))
.header(CONTENT_TYPE, "application/json")
.json(&body),
)
.send()
.await
{
Ok(res) => res,
Err(error) => {
eprintln!(
"chat failed stage=transport category={}",
transport_error_category(&error)
);
return Err(provider());
}
};
if !res.status().is_success() {
eprintln!("chat failed stage=http status={}", res.status());
return Err(provider());
}
match res.json::<ProviderReply>().await {
Ok(parsed) => Ok(parsed),
Err(error) => {
eprintln!("chat failed stage=decode error={error}");
Err(provider())
}
}
}
async fn transcribe(
state: &AppState,
bytes: Vec<u8>,
format: &str,
language: &str,
) -> ApiResult<String> {
eprintln!("STT started bytes={}", bytes.len());
for attempt in 0..2 {
let model = state.config.stt_model.as_ref().ok_or_else(bad)?;
let part = reqwest::multipart::Part::bytes(bytes.clone())
.file_name(format!("audio.{format}"))
.mime_str(&format!("audio/{format}"))
.map_err(|_| bad())?;
let form = reqwest::multipart::Form::new()
.text("model", model.clone())
.text("language", language.to_string())
.part("file", part);
let result = request_headers(
&state.config,
state
.http
.post(format!("{}/audio/transcriptions", state.config.base_url))
.multipart(form),
)
.send()
.await;
match result {
Ok(r) if r.status().is_success() => {
let v: Value = match r.json().await {
Ok(v) => v,
Err(_) => {
eprintln!("STT failed");
return Err(provider());
}
};
let Some(text) = v
.get("text")
.and_then(Value::as_str)
.map(str::trim)
.filter(|s| !s.is_empty())
else {
eprintln!("STT failed");
return Err(provider());
};
eprintln!("STT completed attempt={}", attempt + 1);
return Ok(text.into());
}
Ok(r) if !retry_status(r.status()) => {
eprintln!("STT failed");
return Err(provider());
}
_ if attempt == 1 => {
eprintln!("STT failed");
return Err(provider());
}
_ => tokio::time::sleep(Duration::from_millis(150)).await,
}
}
unreachable!()
}
fn retry_status(s: StatusCode) -> bool {
matches!(s.as_u16(), 408 | 429 | 500 | 502 | 503 | 504)
}
fn transport_error_category(error: &reqwest::Error) -> &'static str {
if error.is_timeout() {
"timeout"
} else if error.is_connect() {
"connect"
} else if error.is_body() {
"body"
} else if error.is_decode() {
"decode"
} else if error.is_request() {
"request"
} else {
"other"
}
}
async fn speak(state: &AppState, text: &str) -> ApiResult<(Vec<u8>, String)> {
eprintln!("TTS started characters={}", text.chars().count());
for attempt in 0..2 {
let request = match tts_request(&state.http, &state.config, text) {
Ok(request) => request,
Err(_) => {
eprintln!("TTS failed attempt={} serialization", attempt + 1);
return Err(provider());
}
};
let result = request.send().await;
match result {
Ok(r) if r.status().is_success() => {
let mime = r
.headers()
.get(header::CONTENT_TYPE)
.and_then(|x| x.to_str().ok())
.unwrap_or("audio/mpeg")
.to_string();
let b = match r.bytes().await {
Ok(bytes) => bytes.to_vec(),
Err(error) => {
eprintln!(
"TTS failed attempt={} transport={}",
attempt + 1,
transport_error_category(&error)
);
return Err(provider());
}
};
if !b.is_empty() {
eprintln!("TTS completed attempt={} bytes={}", attempt + 1, b.len());
return Ok((b, mime));
}
eprintln!("TTS failed attempt={} empty_response", attempt + 1);
return Err(provider());
}
Ok(r) if !retry_status(r.status()) => {
eprintln!("TTS failed attempt={} status={}", attempt + 1, r.status());
return Err(provider());
}
Ok(r) if attempt == 1 => {
eprintln!("TTS failed attempt={} status={}", attempt + 1, r.status());
return Err(provider());
}
Ok(r) => {
eprintln!("TTS retrying attempt={} status={}", attempt + 1, r.status());
tokio::time::sleep(Duration::from_millis(150)).await;
}
Err(error) if attempt == 1 => {
eprintln!(
"TTS failed attempt={} transport={}",
attempt + 1,
transport_error_category(&error)
);
return Err(provider());
}
Err(error) => {
eprintln!(
"TTS retrying attempt={} transport={}",
attempt + 1,
transport_error_category(&error)
);
tokio::time::sleep(Duration::from_millis(150)).await;
}
}
}
unreachable!()
}
// orchestration loops
async fn schedule_loop(state: Arc<AppState>, sid: String) {
let session = { state.sessions.lock().await.get(&sid).cloned() };
let Some(session) = session else { return };
let wait = {
let mut s = session.lock().await;
if s.debounce_scheduled || s.model_running {
return;
}
s.debounce_scheduled = true;
state.config.batch
};
tokio::spawn(async move {
tokio::time::sleep(wait).await;
run_model(state, sid).await;
});
}
async fn run_model(state: Arc<AppState>, sid: String) {
let session = { state.sessions.lock().await.get(&sid).cloned() };
let Some(session) = session else { return };
let (payload, revision, start_delivery) = {
let mut s = session.lock().await;
s.debounce_scheduled = false;
if s.model_running || s.status != SessionStatus::Active {
return;
}
let mut pending = Vec::new();
while let Some(message) = s.pending_messages.pop_front() {
pending.push(message);
}
if pending.is_empty() {
return;
}
for message in pending {
if history_room(&mut s, state.config.max_history, 1).is_err() {
return;
}
s.messages.push(message);
}
s.model_running = true;
let transition = emit(&mut s, None, ServerEvent::new("ai.thinking", json!({})));
(
s.messages.clone(),
s.input_revision,
transition.start_delivery,
)
};
if start_delivery {
trigger_delivery(state.clone(), sid.clone());
}
let result = chat(&state, payload, tool_schema(), true, &sid)
.await
.and_then(|reply| parse_tool(reply).map_err(|_| chat_failure()));
let mut start = false;
let mut syntheses = Vec::new();
{
let mut s = session.lock().await;
s.model_running = false;
// An interrupt or accepted input after this turn was drained invalidates its reply.
// The newer input remains pending and is scheduled below.
if !model_result_is_current(&s, revision) {
if !s.pending_messages.is_empty() {
s.debounce_scheduled = true;
schedule_model_after(state.clone(), sid.clone(), state.config.batch);
}
return;
}
match result.and_then(|(call, args)| {
validate_tool(&s, args)
.map(|x| (call, x))
.map_err(|_| chat_failure())
}) {
Ok((call, dispatches)) => {
if history_room(&mut s, state.config.max_history, 2).is_ok() {
let call_id = call.id.clone();
s.messages.push(ChatMessage {
role: "assistant".into(),
name: None,
content: None,
tool_calls: Some(vec![call]),
tool_call_id: None,
});
let mut count = 0;
for d in dispatches {
let delivery_id = Uuid::new_v4().to_string();
let (participant_id, participant_api_name, acks) = if d.reply_to == "all" {
let ids: HashSet<_> = s
.participants
.values()
.filter(|p| {
p.onboarding == Onboarding::Complete
&& s.sockets.contains_key(&p.id)
})
.map(|p| p.id.clone())
.collect();
("all".to_string(), "all".to_string(), ids)
} else {
let Some(p) =
s.participants.values().find(|p| p.api_name == d.reply_to)
else {
continue;
};
let mut ids = HashSet::new();
ids.insert(p.id.clone());
(p.id.clone(), p.api_name.clone(), ids)
};
s.delivery_acks.insert(delivery_id.clone(), acks);
s.deliveries.push_back(Delivery {
id: delivery_id.clone(),
participant_id,
participant_api_name,
text: d.text,
audio: None,
mime: "audio/mpeg".into(),
next_speaker: d.next_speaker,
ready: state.config.tts.is_none(),
});
if state.config.tts.is_some() {
syntheses.push((
delivery_id,
s.deliveries.back().expect("queued delivery").text.clone(),
));
}
count += 1;
}
s.messages.push(ChatMessage {
role: "tool".into(),
name: None,
content: Some(ChatContent::Text(
json!({"status":"success","queued":count}).to_string(),
)),
tool_calls: None,
tool_call_id: Some(call_id),
});
start = s.active_delivery.is_none() && next_delivery_ready(&s);
}
}
Err(_) => {
start |= emit(
&mut s,
None,
ServerEvent::new("error", json!({"code":"provider_failure"})),
)
.start_delivery;
}
}
if !s.pending_messages.is_empty() {
s.debounce_scheduled = true;
let st = state.clone();
let id = sid.clone();
let delay = state.config.batch;
schedule_model_after(st, id, delay);
}
}
for (delivery_id, text) in syntheses {
tokio::spawn(synthesize_delivery(
state.clone(),
sid.clone(),
delivery_id,
text,
));
}
if start {
trigger_delivery(state, sid);
}
}
async fn synthesize_delivery(state: Arc<AppState>, sid: String, did: String, text: String) {
let spoken = speak(&state, &text).await.ok();
let session = { state.sessions.lock().await.get(&sid).cloned() };
let Some(session) = session else { return };
let start = {
let mut s = session.lock().await;
let Some(delivery) = s.deliveries.iter_mut().find(|d| d.id == did) else {
return;
};
if let Some((audio, mime)) = spoken {
delivery.audio = Some(audio);
delivery.mime = mime;
}
delivery.ready = true;
s.active_delivery.is_none() && s.deliveries.front().is_some_and(|d| d.id == did)
};
if start {
trigger_delivery(state, sid);
}
}
async fn start_delivery(state: Arc<AppState>, sid: String) {
let session = { state.sessions.lock().await.get(&sid).cloned() };
let Some(session) = session else { return };
let (delivery, advanced) = {
let mut s = session.lock().await;
if s.active_delivery.is_some() {
return;
}
let Some(d) = s.deliveries.front() else {
return;
};
if !d.ready {
return;
}
let d = s.deliveries.pop_front().expect("front checked");
let delivery_id = d.id.clone();
emit(
&mut s,
None,
ServerEvent::new(
"delivery.preparing",
json!({"target":d.participant_api_name}),
),
);
s.active_delivery = Some(d);
s.last_activity_at = Instant::now();
(delivery_id, complete_unacknowledgeable_delivery(&mut s))
};
if advanced {
trigger_delivery(state, sid);
return;
}
let state2 = state.clone();
tokio::spawn(async move {
let (text, has_audio) = {
let s = session.lock().await;
let Some(d) = s.active_delivery.as_ref().filter(|d| d.id == delivery) else {
return;
};
if !active_delivery_target_connected(&s, &delivery) {
drop(s);
complete_delivery(state2, sid, delivery).await;
return;
}
(d.text.clone(), d.audio.is_some())
};
let mut s = session.lock().await;
let recipients = s.delivery_acks.get(&delivery).cloned().unwrap_or_default();
let mut advanced = false;
for recipient in recipients {
let outcome = emit(
&mut s,
Some(&recipient),
ServerEvent::new(
"delivery.started",
json!({"delivery_id":delivery,"text":text,"audio":has_audio}),
),
);
advanced |= outcome.start_delivery;
}
if advanced {
drop(s);
trigger_delivery(state2, sid);
return;
}
if !active_delivery_target_connected(&s, &delivery) {
drop(s);
complete_delivery(state2, sid, delivery).await;
return;
}
drop(s);
let timeout = state2.config.delivery_timeout;
tokio::time::sleep(timeout).await;
complete_delivery(state2, sid, delivery).await;
});
}
fn trigger_delivery(state: Arc<AppState>, sid: String) {
tokio::spawn(async move {
start_delivery(state, sid).await;
});
}
fn schedule_model_after(state: Arc<AppState>, sid: String, delay: Duration) {
tokio::spawn(async move {
tokio::time::sleep(delay).await;
run_model(state, sid).await;
});
}
async fn complete_delivery(state: Arc<AppState>, sid: String, did: String) {
let session = { state.sessions.lock().await.get(&sid).cloned() };
let Some(session) = session else { return };
let start;
{
let mut s = session.lock().await;
let Some(active) = s.active_delivery.as_ref() else {
return;
};
if active.id != did {
return;
}
let active = s.active_delivery.take().unwrap();
s.delivery_acks.remove(&did);
s.last_activity_at = Instant::now();
emit(
&mut s,
None,
ServerEvent::new("delivery.completed", json!({"delivery_id":did})),
);
if let Some(next) = active.next_speaker {
s.floor_holder = s
.participants
.values()
.find(|p| p.api_name == next && s.sockets.contains_key(&p.id))
.map(|p| p.id.clone());
let floor_holder = s
.floor_holder
.as_ref()
.and_then(|id| s.participants.get(id))
.map(|p| p.api_name.clone());
emit(
&mut s,
None,
ServerEvent::new("floor.changed", json!({"floor_holder":floor_holder})),
);
}
start = !s.deliveries.is_empty();
}
if start {
trigger_delivery(state, sid);
}
}
// HTTP/WS handlers
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct LanguageRequest {
browser_language: Option<String>,
}
#[derive(Serialize)]
struct Created {
session_id: String,
participant_token: String,
api_name: String,
}
#[derive(Serialize)]
struct OkResponse {
ok: bool,
}
async fn page(State(st): State<Arc<AppState>>) -> Html<String> {
Html(
PAGE.replace(
"__VOICE_MODE__",
if st.config.tts.is_some() {
"true"
} else {
"false"
},
)
.replace("__MAX_AUDIO_BYTES__", &st.config.max_audio.to_string()),
)
}
async fn health() -> &'static str {
"ok"
}
fn start_onboarding_prompt(
state: Arc<AppState>,
sid: String,
participant_id: String,
revision: u64,
) {
tokio::spawn(async move {
let session = { state.sessions.lock().await.get(&sid).cloned() };
let Some(session) = session else { return };
let text = {
let s = session.lock().await;
let Some(p) = s.participants.get(&participant_id) else {
return;
};
let lang = if p.browser_language.is_empty() {
if s.language.is_empty() {
state.config.default_language.clone()
} else {
s.language.clone()
}
} else {
p.browser_language.clone()
};
language_template(&lang).to_string()
};
let audio = if state.config.tts.is_some() {
speak(&state, &text).await.ok()
} else {
None
};
let mut s = session.lock().await;
if !s
.participants
.get(&participant_id)
.is_some_and(|p| p.onboarding == Onboarding::Required)
|| s.onboarding_revision != revision
{
return;
}
let id = Uuid::new_v4().to_string();
let (bytes, mime) = audio.map_or((None, "audio/mpeg".to_string()), |(b, m)| (Some(b), m));
s.onboarding_audio = Some(OnboardingAudio {
id: id.clone(),
bytes,
mime,
text: text.clone(),
});
let has_audio = s
.onboarding_audio
.as_ref()
.is_some_and(|prompt| prompt.bytes.is_some());
emit(
&mut s,
Some(&participant_id),
ServerEvent::new(
"onboarding.prompt",
json!({"delivery_id":id,"text":text,"audio":has_audio}),
),
);
});
}
async fn create(State(st): State<Arc<AppState>>, request: Request) -> ApiResult<Json<Created>> {
let req: LanguageRequest = json_body(request).await?;
let mut map = st.sessions.lock().await;
if map.len() >= st.config.max_sessions {
let ttl = st.config.ttl;
map.retain(|id, session| {
let retain = session
.try_lock()
.map(|s| !session_expired(&s, ttl))
.unwrap_or(true);
if !retain {
eprintln!("session expired id={id}");
}
retain
});
}
if map.len() >= st.config.max_sessions {
return Err(ApiError(
"server_busy",
"The request could not be completed.",
StatusCode::SERVICE_UNAVAILABLE,
));
}
let id = Uuid::new_v4().to_string();
let mut s = now_session(id.clone());
s.onboarding_messages.push(ChatMessage {
role: "system".into(),
name: None,
content: Some(ChatContent::Text(private_system_prompt(&st.config))),
tool_calls: None,
tool_call_id: None,
});
let p = add_participant(
&mut s,
req.browser_language
.unwrap_or_else(|| st.config.default_language.clone()),
true,
);
map.insert(id.clone(), Arc::new(Mutex::new(s)));
drop(map);
eprintln!("private conversation started id={id}");
start_onboarding_prompt(st.clone(), id.clone(), p.id.clone(), 1);
Ok(Json(Created {
session_id: id,
participant_token: p.token,
api_name: p.api_name,
}))
}
async fn join(
State(st): State<Arc<AppState>>,
Path(id): Path<String>,
request: Request,
) -> ApiResult<Json<Created>> {
let req: LanguageRequest = json_body(request).await?;
let session = st.sessions.lock().await.get(&id).cloned().ok_or(ApiError(
"session_not_found",
"The request could not be completed.",
StatusCode::NOT_FOUND,
))?;
let mut s = session.lock().await;
if s.status != SessionStatus::Active {
return Err(ApiError(
"session_not_ready",
"The request could not be completed.",
StatusCode::CONFLICT,
));
}
if s.participants.len() >= st.config.max_participants {
return Err(ApiError(
"session_full",
"The request could not be completed.",
StatusCode::CONFLICT,
));
}
let p = add_participant(
&mut s,
req.browser_language
.unwrap_or_else(|| st.config.default_language.clone()),
false,
);
s.last_activity_at = Instant::now();
s.pending_messages.push_back(event_message(format!(
"Participant joined: {}.",
p.api_name
)));
let completion = emit(
&mut s,
None,
ServerEvent::new("participant.joined", json!({"api_name":p.api_name})),
);
drop(s);
if completion.start_delivery {
trigger_delivery(st.clone(), id.clone());
}
schedule_loop(st.clone(), id.clone()).await;
eprintln!("participant joined session={id} participant={}", p.id);
Ok(Json(Created {
session_id: id,
participant_token: p.token,
api_name: p.api_name,
}))
}
async fn bootstrap_creator(
st: Arc<AppState>,
id: String,
token: String,
user: ChatMessage,
) -> ApiResult<Json<Value>> {
let session = lookup_session(&st, &id).await?;
let (messages, revision) = {
let mut s = session.lock().await;
let participant = require_auth(&s, &token)?;
if !participant.creator
|| participant.onboarding != Onboarding::Required
|| s.status != SessionStatus::Onboarding
{
return Err(bad());
}
let mut messages = s.onboarding_messages.clone();
messages.push(user.clone());
s.onboarding_revision += 1;
(messages, s.onboarding_revision)
};
let schema = json!({"type":"function","function":{"name":"start_group_session","description":"Create a multi-user session only after explicit group intent with a concrete purpose. session_prompt must expand that specific purpose into instructions for this group.","parameters":{"type":"object","additionalProperties":false,"properties":{"display_name":{"type":"string","minLength":1,"maxLength":80},"language":{"type":"string","minLength":2,"maxLength":16},"session_prompt":{"type":"string","minLength":20,"maxLength":8000}},"required":["session_prompt"]}}});
let response = chat(&st, messages, schema, false, &id).await?;
let choice = response
.choices
.into_iter()
.next()
.ok_or_else(chat_failure)?;
if let Some(mut calls) = choice.message.tool_calls {
if choice.finish_reason.as_deref() != Some("tool_calls")
|| choice.message.role != "assistant"
|| choice.message.content.is_some()
|| calls.len() != 1
{
return Err(chat_failure());
}
let call = calls.pop().ok_or_else(chat_failure)?;
if call.kind != "function"
|| call.id.is_empty()
|| call.function.name != "start_group_session"
{
return Err(chat_failure());
}
let args: StartGroupSessionArgs =
serde_json::from_str(&call.function.arguments).map_err(|_| chat_failure())?;
let display = args
.display_name
.as_deref()
.map(str::trim)
.filter(|value| !value.is_empty() && value.chars().count() <= 80)
.map(str::to_string);
if args.display_name.is_some() && display.is_none() {
return Err(chat_failure());
}
let language = args
.language
.as_deref()
.map(str::trim)
.filter(|value| (2..=16).contains(&value.len()))
.map(str::to_string);
if args.language.is_some() && language.is_none() {
return Err(chat_failure());
}
let session_prompt = args.session_prompt.trim();
if !(20..=8000).contains(&session_prompt.chars().count()) {
return Err(chat_failure());
}
let session_prompt = session_prompt.to_string();
let mut s = session.lock().await;
let pid = require_auth(&s, &token)?.id.clone();
if !s.participants[&pid].creator
|| s.participants[&pid].onboarding != Onboarding::Required
|| s.status != SessionStatus::Onboarding
|| s.onboarding_revision != revision
{
return Err(ApiError(
"stale_onboarding",
"The request could not be completed.",
StatusCode::CONFLICT,
));
}
let participant_api = {
let participant = s.participants.get_mut(&pid).ok_or_else(unauthorized)?;
participant.display_name = display.clone();
participant.onboarding = Onboarding::Complete;
participant.api_name.clone()
};
s.language = language.unwrap_or_else(|| st.config.default_language.clone());
let system_prompt =
group_system_prompt(&session_prompt, &st.config.group_session_protocol_prompt);
s.messages = vec![ChatMessage {
role: "system".into(),
name: None,
content: Some(ChatContent::Text(system_prompt)),
tool_calls: None,
tool_call_id: None,
}];
s.pending_messages.push_back(event_message(format!(
"Session started with {}. Begin the conversation.",
participant_api
)));
s.status = SessionStatus::Active;
s.onboarding_messages.clear();
s.onboarding_audio = None;
s.last_activity_at = Instant::now();
emit(
&mut s,
None,
ServerEvent::new("participant.named", json!({"api_name":participant_api})),
);
eprintln!("session created id={id}");
drop(s);
schedule_loop(st, id.clone()).await;
return Ok(Json(
json!({"ok":true,"session_created":true,"session_id":id}),
));
}
let answer = choice
.message
.content
.and_then(ChatContent::text)
.map(|content| content.trim().to_string())
.filter(|content| !content.is_empty())
.ok_or_else(chat_failure)?;
let audio = if st.config.tts.is_some() {
speak(&st, &answer).await.ok()
} else {
None
};
let mut s = session.lock().await;
let pid = require_auth(&s, &token)?.id.clone();
if s.status != SessionStatus::Onboarding || s.onboarding_revision != revision {
return Err(ApiError(
"stale_onboarding",
"The request could not be completed.",
StatusCode::CONFLICT,
));
}
history_room_for_draft(&s.onboarding_messages, st.config.max_history, 2)?;
s.onboarding_messages.push(user);
s.onboarding_messages.push(ChatMessage {
role: "assistant".into(),
name: None,
content: Some(ChatContent::Text(answer.clone())),
tool_calls: None,
tool_call_id: None,
});
let delivery_id = Uuid::new_v4().to_string();
let (bytes, mime) = audio.map_or((None, "audio/mpeg".to_string()), |(bytes, mime)| {
(Some(bytes), mime)
});
s.onboarding_audio = Some(OnboardingAudio {
id: delivery_id.clone(),
bytes,
mime,
text: answer.clone(),
});
let has_audio = s
.onboarding_audio
.as_ref()
.is_some_and(|prompt| prompt.bytes.is_some());
emit(
&mut s,
Some(&pid),
ServerEvent::new(
"onboarding.prompt",
json!({"delivery_id":delivery_id,"text":answer,"audio":has_audio}),
),
);
s.last_activity_at = Instant::now();
Ok(Json(json!({"ok":true,"session_created":false})))
}
fn history_room_for_draft(messages: &[ChatMessage], max: usize, count: usize) -> ApiResult<()> {
if messages.len().saturating_add(count) > max {
return Err(ApiError(
"history_limit",
"The request could not be completed.",
StatusCode::CONFLICT,
));
}
Ok(())
}
fn bearer(h: &HeaderMap) -> Option<&str> {
h.get(header::AUTHORIZATION)?
.to_str()
.ok()?
.strip_prefix("Bearer ")
}
async fn session_view(
State(st): State<Arc<AppState>>,
Path(id): Path<String>,
headers: HeaderMap,
) -> ApiResult<Json<Value>> {
let token = bearer(&headers).ok_or_else(unauthorized)?;
let session = lookup_session(&st, &id).await?;
let s = session.lock().await;
require_auth(&s, token)?;
Ok(Json(public_view(&s)))
}
async fn bootstrap(
State(st): State<Arc<AppState>>,
Path(id): Path<String>,
request: Request,
) -> ApiResult<Json<Value>> {
let token = bearer(request.headers())
.ok_or(ApiError(
"unauthorized",
"The request could not be completed.",
StatusCode::UNAUTHORIZED,
))?
.to_string();
let authorized_session = lookup_session(&st, &id).await?;
let creator_api_name = {
let s = authorized_session.lock().await;
let p = require_auth(&s, &token)?;
if !p.creator
|| p.onboarding != Onboarding::Required
|| s.status != SessionStatus::Onboarding
{
return Err(bad());
}
p.api_name.clone()
};
let content_type = request
.headers()
.get(header::CONTENT_TYPE)
.and_then(|v| v.to_str().ok())
.unwrap_or("");
let user = if content_type.starts_with("application/json") {
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct BootstrapText {
text: String,
}
let text = json_body::<BootstrapText>(request).await?.text;
let text = limited_text(text)?;
ChatMessage {
role: "user".into(),
name: None,
content: Some(ChatContent::Text(text)),
tool_calls: None,
tool_call_id: None,
}
} else {
let (audio, format) = browser_audio(request, st.config.max_audio).await?;
let session = st.sessions.lock().await.get(&id).cloned().ok_or(ApiError(
"session_not_found",
"The request could not be completed.",
StatusCode::NOT_FOUND,
))?;
let language = {
let s = session.lock().await;
let p = require_auth(&s, &token)?;
if !p.browser_language.is_empty() {
p.browser_language.clone()
} else if !s.language.is_empty() {
s.language.clone()
} else {
st.config.default_language.clone()
}
};
if st.config.stt_model.is_some() {
ChatMessage {
role: "user".into(),
name: None,
content: Some(ChatContent::Text(
transcribe(&st, audio, format.as_str(), &language).await?,
)),
tool_calls: None,
tool_call_id: None,
}
} else {
ChatMessage {
role: "user".into(),
name: Some(creator_api_name.clone()),
content: Some(ChatContent::Parts(vec![
ContentPart::Text {
text: format!("{} sent audio during private onboarding.", creator_api_name),
},
ContentPart::InputAudio {
input_audio: AudioPart {
data: STANDARD.encode(audio),
format,
},
},
])),
tool_calls: None,
tool_call_id: None,
}
}
};
let session = st.sessions.lock().await.get(&id).cloned().ok_or(ApiError(
"session_not_found",
"The request could not be completed.",
StatusCode::NOT_FOUND,
))?;
{
let s = session.lock().await;
let p = require_auth(&s, &token)?;
if !p.creator || p.onboarding != Onboarding::Required {
return Err(bad());
}
}
bootstrap_creator(st, id, token, user).await
}
async fn input_text(
State(st): State<Arc<AppState>>,
Path(id): Path<String>,
headers: HeaderMap,
request: Request,
) -> ApiResult<Json<Value>> {
let token = bearer(&headers).ok_or(ApiError(
"unauthorized",
"The request could not be completed.",
StatusCode::UNAUTHORIZED,
))?;
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct TextRequest {
text: String,
}
let v: TextRequest = json_body(request).await?;
let text = limited_text(v.text)?;
accept_input(st, id, token.to_string(), text).await?;
Ok(Json(json!({"ok":true})))
}
async fn interrupt(
State(st): State<Arc<AppState>>,
Path(id): Path<String>,
headers: HeaderMap,
) -> ApiResult<Json<OkResponse>> {
let token = bearer(&headers).ok_or_else(unauthorized)?;
let session = lookup_session(&st, &id).await?;
let mut s = session.lock().await;
let participant_id = require_auth(&s, token)?.id.clone();
let participant = s
.participants
.get(&participant_id)
.ok_or_else(unauthorized)?;
if participant.onboarding != Onboarding::Complete || s.status != SessionStatus::Active {
return Err(ApiError(
"onboarding_required",
"The request could not be completed.",
StatusCode::CONFLICT,
));
}
s.input_revision += 1;
interrupt_deliveries(&mut s, &participant_id);
s.last_activity_at = Instant::now();
Ok(Json(OkResponse { ok: true }))
}
async fn input_audio(
State(st): State<Arc<AppState>>,
Path(id): Path<String>,
request: Request,
) -> ApiResult<Json<Value>> {
let token = bearer(request.headers())
.ok_or(ApiError(
"unauthorized",
"The request could not be completed.",
StatusCode::UNAUTHORIZED,
))?
.to_string();
let session = st.sessions.lock().await.get(&id).cloned().ok_or(ApiError(
"session_not_found",
"The request could not be completed.",
StatusCode::NOT_FOUND,
))?;
let language = {
let s = session.lock().await;
let participant = require_auth(&s, &token)?;
if participant.onboarding != Onboarding::Complete {
return Err(ApiError(
"onboarding_required",
"The request could not be completed.",
StatusCode::CONFLICT,
));
}
if s.status != SessionStatus::Active {
return Err(bad());
}
s.language.clone()
};
let (data, format) = browser_audio(request, st.config.max_audio).await?;
if st.config.stt_model.is_some() {
let text = transcribe(&st, data, format.as_str(), &language).await?;
accept_input(st, id, token, text).await?;
} else {
accept_audio_input(st, id, token, data, format).await?;
}
Ok(Json(json!({"ok":true})))
}
async fn accept_audio_input(
st: Arc<AppState>,
id: String,
token: String,
data: Vec<u8>,
format: AudioFormat,
) -> ApiResult<()> {
let session = lookup_session(&st, &id).await?;
{
let mut s = session.lock().await;
let participant_id = require_auth(&s, &token)?.id.clone();
let (api_name, onboarding) = s
.participants
.get(&participant_id)
.map(|p| (p.api_name.clone(), p.onboarding == Onboarding::Complete))
.ok_or_else(unauthorized)?;
if !onboarding || s.status != SessionStatus::Active {
return Err(bad());
}
s.pending_messages.push_back(ChatMessage {
role: "user".into(),
name: Some(api_name.clone()),
content: Some(ChatContent::Parts(vec![
ContentPart::Text {
text: format!("{} sent audio.", api_name),
},
ContentPart::InputAudio {
input_audio: AudioPart {
data: STANDARD.encode(data),
format,
},
},
])),
tool_calls: None,
tool_call_id: None,
});
s.input_revision += 1;
s.last_activity_at = Instant::now();
interrupt_deliveries(&mut s, &participant_id);
emit(
&mut s,
None,
ServerEvent::new("input.accepted", json!({"api_name":api_name})),
);
}
schedule_loop(st, id).await;
Ok(())
}
async fn accept_input(st: Arc<AppState>, id: String, token: String, text: String) -> ApiResult<()> {
let session = st.sessions.lock().await.get(&id).cloned().ok_or(ApiError(
"session_not_found",
"The request could not be completed.",
StatusCode::NOT_FOUND,
))?;
{
let mut s = session.lock().await;
let participant_id = require_auth(&s, &token)?.id.clone();
if s.participants[&participant_id].onboarding != Onboarding::Complete {
return Err(ApiError(
"onboarding_required",
"The request could not be completed.",
StatusCode::CONFLICT,
));
}
if s.status != SessionStatus::Active {
return Err(bad());
}
let message = user_message(&s.participants[&participant_id], text);
s.pending_messages.push_back(message);
s.input_revision += 1;
interrupt_deliveries(&mut s, &participant_id);
s.last_activity_at = Instant::now();
let api_name = s.participants[&participant_id].api_name.clone();
emit(
&mut s,
None,
ServerEvent::new("input.accepted", json!({"api_name":api_name})),
);
}
schedule_loop(st, id).await;
Ok(())
}
async fn leave(
State(st): State<Arc<AppState>>,
Path(id): Path<String>,
headers: HeaderMap,
) -> ApiResult<Json<Value>> {
let token = bearer(&headers).ok_or(ApiError(
"unauthorized",
"The request could not be completed.",
StatusCode::UNAUTHORIZED,
))?;
let session = st.sessions.lock().await.get(&id).cloned().ok_or(ApiError(
"session_not_found",
"The request could not be completed.",
StatusCode::NOT_FOUND,
))?;
let outcome = {
let mut s = session.lock().await;
let pid = require_auth(&s, token)?.id.clone();
disconnect_participant(&mut s, &pid, st.config.max_history)
};
if let Some(delivery) = outcome.active_delivery {
complete_delivery(st.clone(), id.clone(), delivery).await;
}
if outcome.start_delivery {
trigger_delivery(st.clone(), id.clone());
}
if outcome.queued_event {
schedule_loop(st, id).await;
}
Ok(Json(json!({"ok":true})))
}
async fn audio(
State(st): State<Arc<AppState>>,
Path((id, did)): Path<(String, String)>,
headers: HeaderMap,
) -> ApiResult<Response> {
let token = bearer(&headers).ok_or(ApiError(
"unauthorized",
"The request could not be completed.",
StatusCode::UNAUTHORIZED,
))?;
let session = st.sessions.lock().await.get(&id).cloned().ok_or(ApiError(
"session_not_found",
"The request could not be completed.",
StatusCode::NOT_FOUND,
))?;
let s = session.lock().await;
let p = require_auth(&s, token)?;
let d = s
.active_delivery
.as_ref()
.filter(|d| {
d.id == did
&& s.delivery_acks
.get(&d.id)
.map_or(d.participant_id == p.id, |ids| ids.contains(&p.id))
})
.ok_or(ApiError(
"invalid_request",
"The request could not be completed.",
StatusCode::NOT_FOUND,
))?;
let bytes = d.audio.clone().ok_or(ApiError(
"invalid_request",
"The request could not be completed.",
StatusCode::NOT_FOUND,
))?;
Ok(([(header::CONTENT_TYPE, d.mime.clone())], bytes).into_response())
}
async fn onboarding_audio(
State(st): State<Arc<AppState>>,
Path(id): Path<String>,
headers: HeaderMap,
) -> ApiResult<Response> {
let token = bearer(&headers).ok_or(ApiError(
"unauthorized",
"The request could not be completed.",
StatusCode::UNAUTHORIZED,
))?;
let session = st.sessions.lock().await.get(&id).cloned().ok_or(ApiError(
"session_not_found",
"The request could not be completed.",
StatusCode::NOT_FOUND,
))?;
let s = session.lock().await;
let participant = require_auth(&s, token)?;
if !participant.creator
|| participant.onboarding != Onboarding::Required
|| s.status != SessionStatus::Onboarding
{
return Err(unauthorized());
}
let prompt = s.onboarding_audio.as_ref().ok_or_else(bad)?;
let bytes = prompt.bytes.clone().ok_or_else(bad)?;
Ok(([(header::CONTENT_TYPE, prompt.mime.clone())], bytes).into_response())
}
async fn delivery_complete(
State(st): State<Arc<AppState>>,
Path((id, did)): Path<(String, String)>,
headers: HeaderMap,
) -> ApiResult<Json<Value>> {
let token = bearer(&headers)
.ok_or(ApiError(
"unauthorized",
"The request could not be completed.",
StatusCode::UNAUTHORIZED,
))?
.to_string();
let session = st.sessions.lock().await.get(&id).cloned().ok_or(ApiError(
"session_not_found",
"The request could not be completed.",
StatusCode::NOT_FOUND,
))?;
{
let mut s = session.lock().await;
let participant_id = require_auth(&s, &token)?.id.clone();
if s.active_delivery.as_ref().map(|d| d.id.as_str()) != Some(&did)
|| !s.delivery_acks.get(&did).map_or(
s.active_delivery
.as_ref()
.is_some_and(|d| d.participant_id == participant_id),
|ids| ids.contains(&participant_id),
)
{
return Err(ApiError(
"delivery_conflict",
"The request could not be completed.",
StatusCode::CONFLICT,
));
}
let acks = s.delivery_acks.get_mut(&did).ok_or(ApiError(
"delivery_conflict",
"The request could not be completed.",
StatusCode::CONFLICT,
))?;
acks.remove(&participant_id);
if !acks.is_empty() {
return Ok(Json(json!({"ok":true})));
}
}
complete_delivery(st, id, did).await;
Ok(Json(json!({"ok":true})))
}
async fn ws(
State(st): State<Arc<AppState>>,
Path(id): Path<String>,
upgrade: WebSocketUpgrade,
) -> Response {
upgrade.on_upgrade(move |socket| ws_connected(st, id, socket))
}
async fn ws_connected(st: Arc<AppState>, id: String, mut socket: WebSocket) {
let hello = tokio::time::timeout(Duration::from_secs(10), socket.next())
.await
.ok()
.flatten();
let Some(Ok(Message::Text(text))) = hello else {
return;
};
let token = serde_json::from_str::<Value>(&text).ok().and_then(|v| {
if v.get("type").and_then(Value::as_str) == Some("hello") {
v.get("token").and_then(Value::as_str).map(str::to_string)
} else {
None
}
});
let Some(token) = token else { return };
let session = { st.sessions.lock().await.get(&id).cloned() };
let Some(session) = session else { return };
let pid = {
let s = session.lock().await;
match require_auth(&s, &token) {
Ok(p) => p.id.clone(),
Err(_) => return,
}
};
let connection_id = Uuid::new_v4().to_string();
let (tx, mut rx) = mpsc::channel(32);
let start_delivery = {
let mut s = session.lock().await;
s.sockets.insert(
pid.clone(),
SocketConnection {
id: connection_id.clone(),
tx,
},
);
if let Some(p) = s.participants.get_mut(&pid) {
p.connected = true;
}
s.last_activity_at = Instant::now();
let mut transition = emit(
&mut s,
Some(&pid),
ServerEvent::new("session.ready", json!({})),
);
if let Some(active) = s.active_delivery.as_ref().filter(|delivery| {
s.delivery_acks
.get(&delivery.id)
.is_some_and(|recipients| recipients.contains(&pid))
}) {
let replay = (
active.id.clone(),
active.text.clone(),
active.audio.is_some(),
);
transition.merge(emit(
&mut s,
Some(&pid),
ServerEvent::new(
"delivery.started",
json!({"delivery_id":replay.0,"text":replay.1,"audio":replay.2}),
),
));
}
if s.participants[&pid].onboarding == Onboarding::Required
&& let Some(prompt) = s.onboarding_audio.as_ref()
{
let prompt = (
prompt.id.clone(),
prompt.text.clone(),
prompt.bytes.is_some(),
);
transition.merge(emit(
&mut s,
Some(&pid),
ServerEvent::new(
"onboarding.prompt",
json!({"delivery_id":prompt.0,"text":prompt.1,"audio":prompt.2}),
),
));
}
transition.start_delivery
};
if start_delivery {
trigger_delivery(st.clone(), id.clone());
}
while let Some(event) = rx.recv().await {
if socket
.send(Message::Text(match serde_json::to_string(&event) {
Ok(json) => json.into(),
Err(_) => break,
}))
.await
.is_err()
{
break;
}
}
let outcome = {
let mut s = session.lock().await;
if owns_socket(&s, &pid, &connection_id) {
disconnect_participant(&mut s, &pid, st.config.max_history)
} else {
DisconnectOutcome {
active_delivery: None,
start_delivery: false,
queued_event: false,
}
}
};
if let Some(delivery) = outcome.active_delivery {
complete_delivery(st.clone(), id.clone(), delivery).await;
}
if outcome.start_delivery {
trigger_delivery(st.clone(), id.clone());
}
if outcome.queued_event {
schedule_loop(st, id).await;
}
}
// embedded page
const PAGE: &str = r##"<!doctype html>
<html lang="en">
<head>
<meta charset="utf-8"><meta name="viewport" content="width=device-width, initial-scale=1"><title>QQQ</title>
<style>
*{box-sizing:border-box}body{min-height:100vh;margin:0;display:grid;place-items:center;background:#10121b;color:#e9ecff;font:15px system-ui,sans-serif}.stage{text-align:center}.orb{width:150px;height:150px;border:0;border-radius:50%;cursor:pointer;background:radial-gradient(circle at 35% 30%,#a9d3ff,#665cff 58%,#2a2260);box-shadow:0 0 55px #6b5cff;transition:.2s}.orb[data-state="recording"]{background:#ef5b7a;transform:scale(1.08)}.orb[data-state="ai-speaking"]{animation:pulse 1s infinite}.orb[data-state="error"]{background:#9b405b}@keyframes pulse{50%{transform:scale(1.08)}}#participants{min-height:34px;margin:24px 0}.participant{display:inline-grid;place-items:center;width:26px;height:26px;margin:3px;border-radius:50%;background:#4a506d;font-size:10px;font-style:normal}.participant.disconnected{opacity:.32}.participant.floor{outline:2px solid #8cf}.participant.delivery{box-shadow:0 0 0 3px #c9a7ff}#status{min-height:22px;color:#b9c0df}
</style>
</head>
<body><main class="stage"><button id="orb" class="orb" data-state="entry" aria-label="QQQ orb"></button><div id="participants" aria-label="Participants"></div><div id="status" aria-live="polite"></div></main>
<script>
(() => {
const MAX_BYTES=__MAX_AUDIO_BYTES__, MAX_MS=60000, SILENCE_MS=3000;
const query=new URLSearchParams(location.search);
let sid=query.get("session"), token=sid&&sessionStorage.getItem(`qqq:${sid}`);
const audioInput=true,ttsMode=__VOICE_MODE__;let context,stream,recorder,analyser,source,frame,timer,chunks=[],generation=0,state="entry",deliveryTarget,retry,leaving=false,latestGroupText="",currentAudio;
const orb=document.querySelector("#orb"), people=document.querySelector("#participants"), status=document.querySelector("#status");
const set=(next,text="")=>{state=next;orb.dataset.state=next;status.textContent=text};
const auth=()=>({Authorization:`Bearer ${token}`});
async function request(path,options={}){const r=await fetch(path,{...options,headers:{...auth(),...(options.headers||{})}});if(!r.ok)throw Error("request");return r.json()}
function mute(){stream?.getAudioTracks().forEach(track=>track.enabled=false)}
async function unlock(){
if(!audioInput)return;
context ||= new AudioContext();
const resumed=context.resume(), microphone=stream?Promise.resolve(stream):navigator.mediaDevices.getUserMedia({audio:true});
[,stream]=await Promise.all([resumed,microphone]); mute();
}
function mime(){return window.MediaRecorder&&["audio/webm;codecs=opus","audio/webm","audio/ogg;codecs=opus","audio/ogg","audio/mp4"].find(MediaRecorder.isTypeSupported.bind(MediaRecorder))}
function clearAnalysis(){cancelAnimationFrame(frame);clearTimeout(timer);source?.disconnect();analyser?.disconnect();source=analyser=frame=timer=undefined}
function monitor(gen){
source=context.createMediaStreamSource(stream);analyser=context.createAnalyser();analyser.fftSize=512;source.connect(analyser);
const data=new Uint8Array(analyser.fftSize);let quiet=0;
const sample=()=>{if(gen!==generation||recorder?.state!=="recording")return;analyser.getByteTimeDomainData(data);let peak=0;for(const x of data)peak=Math.max(peak,Math.abs(x-128));const now=performance.now();quiet=peak<5?(quiet||now):0;if(quiet&&now-quiet>=SILENCE_MS)stop(gen);else frame=requestAnimationFrame(sample)};
frame=requestAnimationFrame(sample);
}
function promoteGroup(result){
if(!result.session_created)return false;
const promoted=result.session_id;
if(typeof promoted!=="string"||!promoted||promoted!==sid)throw Error("promotion");
sessionStorage.setItem(`qqq:${promoted}`,token);history.replaceState(null,"",`/?session=${encodeURIComponent(promoted)}`);return true;
}
async function textFallback(active=false,prompt="Type a message."){
const text=window.prompt(active?"Type your message.":prompt);
if(!text?.trim()){if(active)set("idle",latestGroupText||"Press the orb to speak or type.");else set("onboarding","Answer when ready.");return}
if(active)latestGroupText="";
const result=await request(`/api/sessions/${sid}/${active?"input/text":"bootstrap/answer"}`,{method:"POST",headers:{"Content-Type":"application/json"},body:JSON.stringify({text:text.trim()})});if(promoteGroup(result)||active)set("waiting","Waiting.");
}
async function begin(kind){
const type=mime();if(!stream||!type)throw Error("media");const gen=++generation;chunks=[];stream.getAudioTracks().forEach(track=>track.enabled=true);
recorder=new MediaRecorder(stream,{mimeType:type});const actualType=recorder.mimeType||type;
recorder.ondataavailable=e=>{if(gen===generation&&e.data.size){chunks.push(e.data);if(chunks.reduce((n,b)=>n+b.size,0)>MAX_BYTES)stop(gen)}};
recorder.onerror=()=>stop(gen,true);recorder.onstop=()=>finish(gen,kind,actualType);recorder.start(250);monitor(gen);timer=setTimeout(()=>stop(gen),MAX_MS);set("recording",kind==="onboarding"?"Speak, then press the orb to finish.":"Listening.");
}
function stop(gen=generation,cancelled=false){if(gen!==generation)return;clearAnalysis();if(cancelled)generation++;mute();if(recorder?.state==="recording")recorder.stop()}
async function finish(gen,kind,type){
clearAnalysis();const clip=new Blob(chunks,{type});chunks=[];recorder=undefined;mute();if(gen!==generation)return;
if(!clip.size||clip.size>MAX_BYTES){if(kind==="onboarding")return textFallback().catch(fail);set("your-turn","Recording was too short.");return}
try{
if(gen!==generation)return;
const form=new FormData();form.append("audio",clip,`voice.${type.includes("ogg")?"ogg":type.includes("mp4")?"m4a":"webm"}`);set("uploading","Sending.");
const result=await request(`/api/sessions/${sid}/${kind==="onboarding"?"bootstrap/answer":"input/audio"}`,{method:"POST",body:form});if(gen===generation&&(kind!=="onboarding"||promoteGroup(result)))set("waiting","Waiting.");
}catch(_){if(kind==="onboarding")textFallback().catch(fail);else if(gen===generation)set("your-turn","Try again.")}
}
async function play(path){const r=await fetch(path,{headers:auth()});if(!r.ok)throw Error("audio");const url=URL.createObjectURL(await r.blob()),audio=currentAudio=new Audio(url);try{await audio.play();await new Promise(done=>{audio.onended=audio.onerror=done})}finally{if(currentAudio===audio)currentAudio=undefined;URL.revokeObjectURL(url)}}
function interrupt(){generation++;currentAudio?.pause();currentAudio=undefined;deliveryTarget=undefined;if(state==="ai-speaking"||state==="waiting")set("idle","Your turn.")}
async function onboarding(event){set("onboarding",event.text||"Type a message.");if(event.audio)try{await play(`/api/sessions/${sid}/onboarding/audio`)}catch(_){status.textContent=event.text||"Answer when ready."}try{await begin("onboarding")}catch(_){textFallback(false,event.text).catch(fail)}}
async function delivery(event){
const deliveryId=event.delivery_id;latestGroupText=event.text||"";deliveryTarget=deliveryId;set("ai-speaking",event.text||"AI is speaking.");
try{if(event.audio)await play(`/api/sessions/${sid}/deliveries/${event.delivery_id}/audio`);else await new Promise(done=>setTimeout(done,Math.max(1200,Math.min(6000,(event.text||"").length*45))))}
catch(_){await new Promise(done=>setTimeout(done,1500))}
finally{if(deliveryTarget!==deliveryId)return;deliveryTarget=undefined;try{await request(`/api/sessions/${sid}/deliveries/${deliveryId}/complete`,{method:"POST"})}catch(_){}refresh()}
}
function render(view){
const target=view.active_delivery?.target;people.replaceChildren(...view.participants.map(p=>{const el=document.createElement("i");el.className=`participant${p.connected?"":" disconnected"}${p.has_floor?" floor":""}${target===p.api_name||deliveryTarget===p.api_name?" delivery":""}`;el.textContent=(p.display_name||p.api_name).trim().charAt(0).toUpperCase();return el}));
if(["recording","uploading","onboarding","error"].includes(state))return;if(!ttsMode&&latestGroupText){set("idle",latestGroupText);return}if(view.ai_state==="thinking")set("ai-thinking","Thinking.");else set("idle","Press the orb to speak or type.");
}
async function refresh(){if(!sid||!token)return;try{render(await request(`/api/sessions/${sid}`))}catch(_){fail()}}
function fail(){set("error","Unable to continue.")}
function connect(){
if(!sid||!token||leaving)return;const protocol=location.protocol==="https:"?"wss:":"ws:";const ws=new WebSocket(`${protocol}//${location.host}/api/sessions/${sid}/ws`);ws.onopen=()=>ws.send(JSON.stringify({type:"hello",token}));
ws.onmessage=e=>{let event;try{event=JSON.parse(e.data)}catch(_){return}if(event.type==="onboarding.prompt")onboarding(event).catch(fail);else if(event.type==="delivery.started")delivery(event);else if(event.type==="delivery.interrupted") {interrupt();refresh()} else if(event.type==="error")fail();else if(["session.ready","participant.joined","participant.named","participant.left","ai.thinking","delivery.preparing","delivery.completed","floor.changed","input.accepted"].includes(event.type))refresh()};
ws.onclose=()=>{if(!leaving&&sid&&token)retry=setTimeout(connect,1000)};
}
async function startPrivateConversation(unlocked){
const x=await request("/api/sessions",{method:"POST",headers:{"Content-Type":"application/json"},body:JSON.stringify({browser_language:navigator.language})});await unlocked;sid=x.session_id;token=x.participant_token;connect();
}
orb.addEventListener("click",async()=>{
if(["uploading","connecting"].includes(state))return;
try{
const unlocked=unlock();
if(!sid){set("connecting","Starting.");await startPrivateConversation(unlocked)}
else if(!token){set("connecting","Joining.");const x=await request(`/api/sessions/${sid}/join`,{method:"POST",headers:{"Content-Type":"application/json"},body:JSON.stringify({browser_language:navigator.language})});await unlocked;token=x.participant_token;sessionStorage.setItem(`qqq:${sid}`,token);connect()}
else {const view=await request(`/api/sessions/${sid}`);await unlocked;if(state==="recording")stop();else {if(state==="ai-speaking"){await request(`/api/sessions/${sid}/interrupt`,{method:"POST"});interrupt()}render(view);if(audioInput)await begin("main");else await textFallback(true)}}
}catch(_){fail()}
});
if(sid&&token){set("connecting","Connecting.");connect();refresh()}setInterval(refresh,12000);
addEventListener("pagehide",()=>{leaving=true;clearTimeout(retry);if(sid&&token)fetch(`/api/sessions/${sid}/leave`,{method:"POST",headers:auth(),keepalive:true}).catch(()=>{})});
})();
</script></body></html>"##;
// main
fn build_router(state: Arc<AppState>) -> Router {
Router::new()
.layer(DefaultBodyLimit::max(MAX_JSON_BYTES))
.route("/", get(page))
.route("/health", get(health))
.route("/api/sessions", post(create))
.route("/api/sessions/{id}/join", post(join))
.route(
"/api/sessions/{id}/bootstrap/answer",
post(bootstrap).layer(DefaultBodyLimit::disable()),
)
.route("/api/sessions/{id}/onboarding/audio", get(onboarding_audio))
.route("/api/sessions/{id}/input/text", post(input_text))
.route("/api/sessions/{id}/interrupt", post(interrupt))
.route(
"/api/sessions/{id}/input/audio",
post(input_audio).layer(DefaultBodyLimit::disable()),
)
.route("/api/sessions/{id}/leave", post(leave))
.route("/api/sessions/{id}/deliveries/{did}/audio", get(audio))
.route(
"/api/sessions/{id}/deliveries/{did}/complete",
post(delivery_complete),
)
.route("/api/sessions/{id}/ws", get(ws))
.route("/api/sessions/{id}", get(session_view))
.with_state(state)
}
#[tokio::main]
async fn main() {
let _ = dotenvy::dotenv();
let config = match Config::load() {
Ok(c) => c,
Err(e) => {
eprintln!("{}", startup_config_error(&e));
std::process::exit(1);
}
};
let bind = config.bind.clone();
let state = Arc::new(AppState {
config,
http: reqwest::Client::builder()
.timeout(Duration::from_secs(45))
.build()
.expect("client"),
sessions: Mutex::new(HashMap::new()),
});
let cleanup = state.clone();
tokio::spawn(async move {
let mut tick = tokio::time::interval(Duration::from_secs(60));
loop {
tick.tick().await;
let snapshot: Vec<_> = cleanup
.sessions
.lock()
.await
.iter()
.map(|(id, s)| (id.clone(), s.clone()))
.collect();
let mut expired = Vec::new();
for (id, session) in snapshot {
if session_expired(&*session.lock().await, cleanup.config.ttl) {
expired.push((id, session));
}
}
let mut map = cleanup.sessions.lock().await;
for (id, candidate) in expired {
let removable = map.get(&id).is_some_and(|current| {
Arc::ptr_eq(current, &candidate)
&& current
.try_lock()
.map(|s| session_expired(&s, cleanup.config.ttl))
.unwrap_or(false)
});
if removable && map.remove(&id).is_some() {
eprintln!("session expired id={id}");
}
}
}
});
let app = build_router(state);
let listener = tokio::net::TcpListener::bind(&bind).await.expect("bind");
eprintln!("server started");
axum::serve(listener, app).await.expect("server");
}
#[cfg(test)]
mod tests {
use super::*;
use axum::body::Body;
use futures_util::{SinkExt, StreamExt};
use tokio_tungstenite::{connect_async, tungstenite::Message as TungsteniteMessage};
use tower::ServiceExt;
#[test]
fn startup_configuration_error_names_the_variable_without_a_value() {
assert_eq!(
startup_config_error("missing OPENAI_API_KEY"),
"configuration error: missing OPENAI_API_KEY; set it in .env or export it"
);
}
#[test]
fn voice_configuration_requires_a_complete_set() {
let neither =
voice_config(String::new(), String::new(), String::new(), "mp3".into()).unwrap();
assert!(neither.0.is_none() && neither.1.is_none());
let both = voice_config("stt".into(), "tts".into(), "voice".into(), "mp3".into()).unwrap();
assert!(both.0.is_some() && both.1.is_some());
assert!(voice_config(String::new(), "tts".into(), String::new(), "mp3".into()).is_err());
}
#[tokio::test]
async fn page_renders_configured_mode() {
let voice = build_router(test_state())
.oneshot(Request::builder().uri("/").body(Body::empty()).unwrap())
.await
.unwrap();
let voice_body = to_bytes(voice.into_body(), usize::MAX).await.unwrap();
assert!(
std::str::from_utf8(&voice_body)
.unwrap()
.contains("const audioInput=true,ttsMode=true")
);
assert!(
std::str::from_utf8(&voice_body)
.unwrap()
.contains("const MAX_BYTES=1")
);
let text_state = test_state();
let mut config = text_state.config.clone();
config.tts = None;
let text = build_router(Arc::new(AppState {
config,
http: reqwest::Client::new(),
sessions: Mutex::new(HashMap::new()),
}))
.oneshot(Request::builder().uri("/").body(Body::empty()).unwrap())
.await
.unwrap();
let text_body = to_bytes(text.into_body(), usize::MAX).await.unwrap();
let text_page = std::str::from_utf8(&text_body).unwrap();
assert!(text_page.contains("const audioInput=true,ttsMode=false"));
assert!(text_page.contains("const MAX_BYTES=1"));
assert!(text_page.contains("latestGroupText"));
assert!(text_page.contains("textFallback(false,event.text"));
}
#[test]
fn private_browser_creation_defers_public_persistence_until_promotion() {
let start = PAGE
.split("async function startPrivateConversation")
.nth(1)
.expect("private creation helper")
.split("orb.addEventListener")
.next()
.expect("private creation helper boundary");
assert!(!start.contains("sessionStorage"));
assert!(!start.contains("history.replaceState"));
let promotion = PAGE
.split("function promoteGroup")
.nth(1)
.expect("promotion helper")
.split("async function textFallback")
.next()
.expect("promotion helper boundary");
assert!(promotion.contains("if(!result.session_created)return false"));
assert!(promotion.contains("promoted!==sid"));
assert!(promotion.contains("sessionStorage.setItem"));
assert!(promotion.contains("history.replaceState"));
}
#[test]
fn configurable_prompts_are_used_for_private_and_group_system_messages() {
let mut config = test_state().config.clone();
config.private_conversation_prompt = "private instructions".into();
config.group_session_protocol_prompt = "group protocol".into();
assert_eq!(private_system_prompt(&config), "private instructions");
assert_eq!(
group_system_prompt("generated purpose", &config.group_session_protocol_prompt),
"generated purpose\n\ngroup protocol"
);
}
fn test_state() -> Arc<AppState> {
Arc::new(AppState {
config: Config {
bind: "127.0.0.1:0".into(),
base_url: "http://localhost".into(),
api_key: "test".into(),
chat_model: "test".into(),
stt_model: Some("test".into()),
tts: Some(TtsConfig {
tts_model: "test".into(),
tts_voice: "test".into(),
tts_format: "mp3".into(),
}),
referer: None,
title: None,
max_sessions: 1,
max_participants: 1,
ttl: Duration::from_secs(1),
max_history: 1,
max_audio: 1,
batch: Duration::from_millis(1),
delivery_timeout: Duration::from_secs(1),
default_language: "en".into(),
private_conversation_prompt: DEFAULT_PRIVATE_CONVERSATION_PROMPT.into(),
group_session_protocol_prompt: DEFAULT_GROUP_SESSION_PROTOCOL_PROMPT.into(),
},
http: reqwest::Client::new(),
sessions: Mutex::new(HashMap::new()),
})
}
fn endpoint_state() -> Arc<AppState> {
let state = test_state();
let mut config = state.config.clone();
config.max_sessions = 10;
config.max_participants = 10;
config.max_history = 512;
config.max_audio = 32;
config.batch = Duration::from_secs(3600);
Arc::new(AppState {
config,
http: reqwest::Client::new(),
sessions: Mutex::new(HashMap::new()),
})
}
async fn seeded_state() -> (Arc<AppState>, String, String, String) {
let state = endpoint_state();
let mut session = ready();
let participant_id = session.participant_order[0].clone();
let token = session.participants[&participant_id].token.clone();
session.floor_holder = Some(participant_id);
let id = session.id.clone();
state
.sessions
.lock()
.await
.insert(id.clone(), Arc::new(Mutex::new(session)));
(state, id, token, "user_1".into())
}
fn request(
method: &str,
uri: String,
token: Option<&str>,
body: Body,
content_type: Option<&str>,
) -> Request {
let mut builder = Request::builder().method(method).uri(uri);
if let Some(token) = token {
builder = builder.header(header::AUTHORIZATION, format!("Bearer {token}"));
}
if let Some(content_type) = content_type {
builder = builder.header(header::CONTENT_TYPE, content_type);
}
builder.body(body).expect("valid test request")
}
async fn response_json(response: Response) -> Value {
let bytes = to_bytes(response.into_body(), usize::MAX)
.await
.expect("test response body");
serde_json::from_slice(&bytes).expect("JSON response")
}
async fn replacement_case(broadcast: bool) {
let (state, sid, first_token, _) = seeded_state().await;
let session = state.sessions.lock().await[&sid].clone();
let second_token = {
let mut s = session.lock().await;
let second = add_participant(&mut s, "en".into(), false);
s.participants
.get_mut(&second.id)
.expect("participant")
.onboarding = Onboarding::Complete;
second.token
};
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("listener");
let address = listener.local_addr().expect("address");
let server = tokio::spawn({
let app = build_router(state.clone());
async move { axum::serve(listener, app).await.expect("server") }
});
let endpoint = format!("ws://{address}/api/sessions/{sid}/ws");
let connect = |token: String| {
let endpoint = endpoint.clone();
async move {
let (mut socket, _) = connect_async(endpoint).await.expect("WebSocket");
socket
.send(TungsteniteMessage::Text(
json!({"type":"hello","token":token}).to_string().into(),
))
.await
.expect("hello");
let ready = tokio::time::timeout(Duration::from_secs(2), socket.next())
.await
.expect("ready timeout")
.expect("open")
.expect("event");
assert_eq!(
serde_json::from_str::<Value>(ready.into_text().expect("text").as_ref())
.expect("JSON")["type"],
"session.ready"
);
socket
}
};
let _old = connect(first_token.clone()).await;
let mut other = connect(second_token.clone()).await;
{
let mut s = session.lock().await;
let recipients = if broadcast {
HashSet::from([
s.participant_order[0].clone(),
s.participant_order[1].clone(),
])
} else {
HashSet::from([s.participant_order[0].clone()])
};
s.delivery_acks.insert("active".into(), recipients);
s.active_delivery = Some(Delivery {
id: "active".into(),
participant_id: if broadcast {
"all"
} else {
&s.participant_order[0]
}
.into(),
participant_api_name: if broadcast { "all" } else { "user_1" }.into(),
text: "exact replay".into(),
audio: Some(vec![7]),
mime: "audio/mpeg".into(),
next_speaker: None,
ready: true,
});
}
let mut replacement = connect(first_token.clone()).await;
let replay = tokio::time::timeout(Duration::from_secs(2), replacement.next())
.await
.expect("replay timeout")
.expect("open")
.expect("event");
let replay: Value =
serde_json::from_str(replay.into_text().expect("text").as_ref()).expect("JSON");
assert_eq!(
replay,
json!({"type":"delivery.started","delivery_id":"active","text":"exact replay","audio":true})
);
assert!(
tokio::time::timeout(Duration::from_millis(150), other.next())
.await
.is_err(),
"replay leaked to another socket"
);
let app = build_router(state.clone());
for token in if broadcast {
vec![&first_token, &second_token]
} else {
vec![&first_token]
} {
let response = app
.clone()
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/deliveries/active/complete"),
Some(token),
Body::empty(),
None,
))
.await
.expect("ack");
assert_eq!(response.status(), StatusCode::OK);
}
assert!(session.lock().await.active_delivery.is_none());
server.abort();
}
#[tokio::test]
async fn targeted_replacement_replays_active_delivery_after_ready() {
replacement_case(false).await;
}
#[tokio::test]
async fn broadcast_replacement_replays_active_delivery_after_ready() {
replacement_case(true).await;
}
#[tokio::test]
async fn unrelated_broadcast_saturation_completes_and_advances_fifo_without_timeout() {
for event in ["participant.left", "floor.changed", "error"] {
let (state, sid, _, _) = seeded_state().await;
let session = state.sessions.lock().await[&sid].clone();
{
let mut s = session.lock().await;
let first = s.participant_order[0].clone();
let second = add_participant(&mut s, "en".into(), false);
s.participants
.get_mut(&second.id)
.expect("participant")
.onboarding = Onboarding::Complete;
let (full_tx, _full_rx) = mpsc::channel(1);
full_tx
.try_send(ServerEvent::new("occupied", json!({})))
.expect("fill socket");
s.sockets.insert(
first.clone(),
SocketConnection {
id: "full".into(),
tx: full_tx,
},
);
let (next_tx, _next_rx) = mpsc::channel(4);
s.sockets.insert(
second.id.clone(),
SocketConnection {
id: "next".into(),
tx: next_tx,
},
);
s.delivery_acks
.insert("active".into(), HashSet::from([first.clone()]));
s.active_delivery = Some(delivery(first, "active", "unused"));
s.deliveries
.push_back(delivery(second.id, "next", "unused"));
assert!(emit(&mut s, None, ServerEvent::new(event, json!({}))).start_delivery);
}
start_delivery(state, sid).await;
assert_eq!(
session
.lock()
.await
.active_delivery
.as_ref()
.map(|delivery| delivery.id.as_str()),
Some("next")
);
}
}
fn multipart(mime: &str, bytes: &[u8]) -> (String, String) {
let boundary = "qqq-test-boundary";
(
format!("multipart/form-data; boundary={boundary}"),
format!(
"--{boundary}\r\nContent-Disposition: form-data; name=\"audio\"; filename=\"voice\"\r\nContent-Type: {mime}\r\n\r\n{}\r\n--{boundary}--\r\n",
String::from_utf8_lossy(bytes)
),
)
}
async fn assert_error(response: Response, status: StatusCode, code: &str) {
assert_eq!(response.status(), status);
assert_eq!(response.headers()[header::CONTENT_TYPE], "application/json");
let body = response_json(response).await;
assert_eq!(
body,
json!({"error":{"code":code,"message":"The request could not be completed."}})
);
}
#[test]
fn router_builds_with_axum_08_path_syntax() {
let _router = build_router(test_state());
}
#[test]
fn tts_request_has_standard_audio_http_shape() {
let state = test_state();
let request = tts_request(&state.http, &state.config, "Hello, world!")
.expect("TTS request serializes")
.build()
.expect("TTS request builds");
assert_eq!(request.method(), reqwest::Method::POST);
assert_eq!(request.url().path(), "/audio/speech");
assert_eq!(
request.headers().get(CONTENT_TYPE).unwrap(),
"application/json"
);
assert_eq!(request.headers().get(ACCEPT).unwrap(), "audio/*");
let body = request
.body()
.and_then(reqwest::Body::as_bytes)
.expect("owned TTS request body");
assert_eq!(
request
.headers()
.get(CONTENT_LENGTH)
.unwrap()
.to_str()
.expect("numeric Content-Length")
.parse::<usize>()
.expect("numeric Content-Length"),
body.len()
);
let payload: TtsRequest = serde_json::from_slice(body).expect("valid TTS JSON");
assert_eq!(payload.model, "test");
assert_eq!(payload.voice, "test");
assert_eq!(payload.input, "Hello, world!");
assert_eq!(payload.response_format, "mp3");
}
#[test]
fn direct_audio_content_serializes_raw_base64_and_required_format() {
let content = ChatContent::Parts(vec![
ContentPart::Text {
text: "user_1 sent audio.".into(),
},
ContentPart::InputAudio {
input_audio: AudioPart {
data: STANDARD.encode([0_u8, 1, 2]),
format: AudioFormat::Wav,
},
},
]);
assert_eq!(
serde_json::to_value(content).expect("serializable direct audio"),
json!([
{"type":"text","text":"user_1 sent audio."},
{"type":"input_audio","input_audio":{"data":"AAEC","format":"wav"}}
])
);
}
#[tokio::test]
async fn wav_input_route_queues_canonical_multimodal_message_without_stt() {
let mut config = endpoint_state().config.clone();
config.stt_model = None;
let state = Arc::new(AppState {
config,
http: reqwest::Client::new(),
sessions: Mutex::new(HashMap::new()),
});
let session = ready();
let sid = session.id.clone();
let participant_id = session.participant_order[0].clone();
let token = session.participants[&participant_id].token.clone();
state
.sessions
.lock()
.await
.insert(sid.clone(), Arc::new(Mutex::new(session)));
let wav = b"RIFFcanonical-wav";
let (content_type, payload) = multipart("audio/wav", wav);
let response = build_router(state.clone())
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/input/audio"),
Some(&token),
Body::from(payload),
Some(&content_type),
))
.await
.expect("audio route response");
assert_eq!(response.status(), StatusCode::OK);
let session = state.sessions.lock().await[&sid].clone();
let session = session.lock().await;
assert_eq!(
serde_json::to_value(session.pending_messages.front().expect("queued audio"))
.expect("message JSON"),
json!({
"role":"user", "name":"user_1", "content":[
{"type":"text","text":"user_1 sent audio."},
{"type":"input_audio","input_audio":{"data":STANDARD.encode(wav),"format":"wav"}}
]
})
);
}
fn ready() -> Session {
let mut s = now_session("s".into());
let p = add_participant(&mut s, "en".into(), true);
let pid = p.id.clone();
let p = s.participants.get_mut(&pid).unwrap();
p.display_name = Some("Ada".into());
p.onboarding = Onboarding::Complete;
test_connect(&mut s, &pid);
s.status = SessionStatus::Active;
s.messages.push(ChatMessage {
role: "system".into(),
name: None,
content: Some("stable".into()),
tool_calls: None,
tool_call_id: None,
});
s
}
fn test_connect(s: &mut Session, id: &str) {
let (tx, rx) = mpsc::channel(32);
// Unit-state fixtures need a live socket identity without a WebSocket task.
std::mem::forget(rx);
s.sockets.insert(
id.into(),
SocketConnection {
id: format!("test-socket-{id}"),
tx,
},
);
s.participants
.get_mut(id)
.expect("test participant")
.connected = true;
}
fn args(to: &str) -> SendMessagesArgs {
SendMessagesArgs {
dispatches: vec![Dispatch {
reply_to: to.into(),
text: "hello".into(),
next_speaker: None,
}],
}
}
fn call(a: &str) -> ToolCall {
ToolCall {
id: "call".into(),
kind: "function".into(),
function: ToolFunction {
name: "send_messages".into(),
arguments: a.into(),
},
}
}
fn delivery(participant_id: String, id: &str, _batch_id: &str) -> Delivery {
Delivery {
id: id.into(),
participant_id,
participant_api_name: "user_1".into(),
text: "x".into(),
audio: None,
mime: "audio/mpeg".into(),
next_speaker: None,
ready: true,
}
}
#[test]
fn system_message_remains_first() {
let s = ready();
assert_eq!(s.messages[0].role, "system");
}
#[test]
fn participant_messages_use_stable_ascii_name() {
let s = ready();
let p = s.participants.values().next().unwrap();
assert_eq!(user_message(p, "x".into()).name, Some("user_1".into()));
}
#[test]
fn content_uses_display_name_text() {
let s = ready();
assert_eq!(
user_message(s.participants.values().next().unwrap(), "hi".into()).content,
Some("Ada: hi".into())
);
}
#[test]
fn assistant_tool_call_remains_in_history() {
let mut s = ready();
s.messages.push(ChatMessage {
role: "assistant".into(),
name: None,
content: None,
tool_calls: Some(vec![call("{}")]),
tool_call_id: None,
});
assert!(s.messages.last().unwrap().tool_calls.is_some());
}
#[test]
fn matching_tool_result_remains_in_history() {
let mut s = ready();
s.messages.push(ChatMessage {
role: "tool".into(),
name: None,
content: Some("{}".into()),
tool_calls: None,
tool_call_id: Some("call".into()),
});
assert_eq!(
s.messages.last().unwrap().tool_call_id.as_deref(),
Some("call")
);
}
#[test]
fn tool_call_result_pairing_never_truncated_independently() {
let mut s = ready();
assert!(history_room(&mut s, 2, 2).is_err());
assert_eq!(s.messages.len(), 1);
}
#[test]
fn participant_join_does_not_rewrite_system_prompt() {
let mut s = ready();
let old = s.messages[0].content.clone();
s.messages
.push(event_message("Participant joined: user_2 (B).".into()));
assert_eq!(s.messages[0].content, old);
}
#[test]
fn join_bootstrap_stores_only_participant_identity() {
let mut s = ready();
let p = add_participant(&mut s, "en".into(), false);
assert!(s.participants[&p.id].display_name.is_none());
}
#[test]
fn bootstrap_messages_never_enter_main_history() {
let s = ready();
assert_eq!(s.messages.len(), 1);
}
#[test]
fn duplicate_display_names_get_unique_api_names() {
let mut s = ready();
let p = add_participant(&mut s, "en".into(), false);
assert_eq!(p.api_name, "user_2");
}
#[test]
fn unknown_reply_to_rejected() {
let s = ready();
assert!(validate_tool(&s, args("nope")).is_err());
}
#[test]
fn all_remains_one_broadcast_dispatch() {
let mut s = ready();
let p = add_participant(&mut s, "en".into(), false);
s.participants.get_mut(&p.id).unwrap().onboarding = Onboarding::Complete;
let v = validate_tool(&s, args("all")).unwrap();
assert_eq!(
v.iter().map(|x| x.reply_to.as_str()).collect::<Vec<_>>(),
vec!["all"]
);
}
#[test]
fn three_participant_dispatch_queue_preserves_recipient_order() {
let mut session = ready();
for name in ["Grace", "Linus"] {
let participant = add_participant(&mut session, "en".into(), false);
let participant_id = participant.id.clone();
let participant = session.participants.get_mut(&participant_id).unwrap();
participant.display_name = Some(name.into());
participant.onboarding = Onboarding::Complete;
test_connect(&mut session, &participant_id);
}
let dispatches = validate_tool(
&session,
SendMessagesArgs {
dispatches: vec![Dispatch {
reply_to: "all".into(),
text: "Please share one idea.".into(),
next_speaker: Some("user_2".into()),
}],
},
)
.expect("valid three-participant dispatch");
assert_eq!(
dispatches
.iter()
.map(|dispatch| dispatch.reply_to.as_str())
.collect::<Vec<_>>(),
["all"]
);
assert_eq!(
dispatches
.iter()
.map(|dispatch| dispatch.next_speaker.as_deref())
.collect::<Vec<_>>(),
[Some("user_2")]
);
assert_eq!(dispatches.len(), 1);
}
#[test]
fn empty_dispatch_rejected() {
let s = ready();
assert!(validate_tool(&s, SendMessagesArgs { dispatches: vec![] }).is_err());
}
#[test]
fn invalid_next_speaker_rejected() {
let s = ready();
let mut a = args("user_1");
a.dispatches[0].next_speaker = Some("bad".into());
assert!(validate_tool(&s, a).is_err());
}
#[test]
fn only_one_model_call_may_run_per_session() {
let mut s = ready();
s.model_running = true;
assert!(s.model_running);
}
#[test]
fn batched_inputs_remain_ordered() {
let mut s = ready();
s.pending_messages.push_back(event_message("one".into()));
s.pending_messages.push_back(event_message("two".into()));
assert_eq!(
s.pending_messages
.pop_front()
.unwrap()
.content
.as_ref()
.and_then(ChatContent::as_text),
Some("one")
);
assert_eq!(
s.pending_messages
.pop_front()
.unwrap()
.content
.as_ref()
.and_then(ChatContent::as_text),
Some("two")
);
}
#[test]
fn messages_during_model_call_trigger_later_call() {
let mut s = ready();
s.model_running = true;
s.pending_messages.push_back(event_message("x".into()));
assert!(!s.pending_messages.is_empty());
}
#[test]
fn only_one_active_delivery_exists() {
let mut s = ready();
s.active_delivery = None;
assert!(s.active_delivery.is_none());
}
#[test]
fn completion_starts_next_delivery() {
let mut s = ready();
s.deliveries.push_back(Delivery {
id: "d".into(),
participant_id: "x".into(),
participant_api_name: "user_1".into(),
text: "x".into(),
audio: None,
mime: "audio/mpeg".into(),
next_speaker: None,
ready: true,
});
assert!(!s.deliveries.is_empty());
}
#[test]
fn disconnected_recipient_timeout_advances_queue() {
let mut s = ready();
let participant_id = s.participant_order[0].clone();
s.participants.get_mut(&participant_id).unwrap().connected = false;
s.active_delivery = Some(delivery(participant_id, "d", "b"));
assert!(!active_delivery_target_connected(&s, "d"));
}
#[test]
fn tts_failure_completes_through_text_fallback() {
let d = Delivery {
id: "d".into(),
participant_id: "p".into(),
participant_api_name: "user_1".into(),
text: "caption".into(),
audio: None,
mime: "audio/mpeg".into(),
next_speaker: None,
ready: true,
};
assert!(d.audio.is_none() && !d.text.is_empty());
}
#[test]
fn fifo_playback_waits_for_the_first_delivery_to_be_ready() {
let mut s = ready();
let participant_id = s.participant_order[0].clone();
let mut first = delivery(participant_id.clone(), "first", "unused");
first.ready = false;
let second = delivery(participant_id, "second", "unused");
s.deliveries.push_back(first);
s.deliveries.push_back(second);
assert!(!next_delivery_ready(&s));
s.deliveries.front_mut().unwrap().ready = true;
assert!(next_delivery_ready(&s));
}
#[test]
fn repeated_disconnect_is_idempotent() {
let mut s = ready();
let id = s.participant_order[0].clone();
let first = disconnect_participant(&mut s, &id, 10);
let second = disconnect_participant(&mut s, &id, 10);
assert!(first.queued_event);
assert!(!second.queued_event);
assert_eq!(s.pending_messages.len(), 1);
}
#[test]
fn stale_socket_does_not_own_replacement_connection() {
let mut session = ready();
let participant_id = session.participant_order[0].clone();
let (old_tx, _old_rx) = mpsc::channel(1);
session.sockets.insert(
participant_id.clone(),
SocketConnection {
id: "old".into(),
tx: old_tx,
},
);
let (new_tx, _new_rx) = mpsc::channel(1);
session.sockets.insert(
participant_id.clone(),
SocketConnection {
id: "new".into(),
tx: new_tx,
},
);
assert!(!owns_socket(&session, &participant_id, "old"));
assert!(owns_socket(&session, &participant_id, "new"));
assert!(session.participants[&participant_id].connected);
}
#[test]
fn expired_sessions_are_removable() {
let mut s = ready();
s.last_activity_at = Instant::now() - Duration::from_secs(2);
assert!(s.last_activity_at.elapsed() > Duration::from_secs(1));
}
#[test]
fn history_limit_closes_without_breaking_tool_pairs() {
let mut s = ready();
assert!(history_room(&mut s, 1, 2).is_err());
assert_eq!(s.status, SessionStatus::Closed);
}
#[test]
fn exact_token_never_in_public_session_view() {
let mut s = ready();
let token = s.participants.values().next().unwrap().token.clone();
s.participants.values_mut().next().unwrap().token = token;
assert!(
!json!({"participants":[{"api_name":"user_1"}]})
.to_string()
.contains(&s.participants.values().next().unwrap().token)
);
}
#[test]
fn tool_arguments_parse_correctly() {
let r = ProviderReply {
choices: vec![Choice {
finish_reason: Some("tool_calls".into()),
message: ProviderMessage {
role: "assistant".into(),
content: None,
tool_calls: Some(vec![call(
r#"{"dispatches":[{"reply_to":"user_1","text":"x","next_speaker":null}]}"#,
)]),
},
}],
};
assert!(parse_tool(r).is_ok());
}
#[test]
fn assistant_ordinary_text_without_tool_rejected() {
let r = ProviderReply {
choices: vec![Choice {
finish_reason: Some("stop".into()),
message: ProviderMessage {
role: "assistant".into(),
content: Some("x".into()),
tool_calls: None,
},
}],
};
assert!(parse_tool(r).is_err());
}
#[test]
fn multiple_tool_calls_rejected() {
let r = ProviderReply {
choices: vec![Choice {
finish_reason: Some("tool_calls".into()),
message: ProviderMessage {
role: "assistant".into(),
content: None,
tool_calls: Some(vec![call("{}"), call("{}")]),
},
}],
};
assert!(parse_tool(r).is_err());
}
#[test]
fn malformed_provider_success_becomes_provider_failure() {
let r = ProviderReply { choices: vec![] };
assert!(parse_tool(r).is_err());
}
#[tokio::test]
async fn provisional_conversation_is_not_joinable_or_an_existing_session() {
let state = endpoint_state();
let app = build_router(state.clone());
let root = app
.clone()
.oneshot(request("GET", "/".into(), None, Body::empty(), None))
.await
.expect("router");
assert_eq!(root.status(), StatusCode::OK);
let health = app
.clone()
.oneshot(request("GET", "/health".into(), None, Body::empty(), None))
.await
.expect("router");
assert_eq!(health.status(), StatusCode::OK);
let created = app
.clone()
.oneshot(request(
"POST",
"/api/sessions".into(),
None,
Body::from(r#"{"browser_language":"en-US"}"#),
Some("application/json"),
))
.await
.expect("router");
assert_eq!(created.status(), StatusCode::OK);
let created = response_json(created).await;
let sid = created["session_id"].as_str().expect("session id");
assert_eq!(created["api_name"], "user_1");
let joined = app
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/join"),
None,
Body::from("{}"),
Some("application/json"),
))
.await
.expect("router");
assert_error(joined, StatusCode::CONFLICT, "session_not_ready").await;
}
#[tokio::test]
async fn active_session_join_is_immediately_anonymous_and_complete() {
let state = endpoint_state();
let session = ready();
let sid = session.id.clone();
state
.sessions
.lock()
.await
.insert(sid.clone(), Arc::new(Mutex::new(session)));
let response = build_router(state.clone())
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/join"),
None,
Body::from(r#"{"browser_language":"en"}"#),
Some("application/json"),
))
.await
.unwrap();
assert_eq!(response.status(), StatusCode::OK);
let joined = response_json(response).await;
let session = state.sessions.lock().await[&sid].clone();
let s = session.lock().await;
let participant = s
.participants
.values()
.find(|p| p.api_name == joined["api_name"])
.unwrap();
assert!(participant.onboarding == Onboarding::Complete);
assert!(participant.display_name.is_none());
assert_eq!(
user_message(participant, "hello".into())
.content
.as_ref()
.and_then(ChatContent::as_text),
Some("user_2: hello")
);
assert_eq!(
s.pending_messages
.back()
.and_then(|message| message.content.as_ref().and_then(ChatContent::as_text)),
Some("Participant joined: user_2.")
);
}
#[tokio::test]
async fn endpoint_view_auth_and_credentials_are_private() {
let (state, sid, token, _) = seeded_state().await;
let app = build_router(state);
let unauthorized = app
.clone()
.oneshot(request(
"GET",
format!("/api/sessions/{sid}"),
None,
Body::empty(),
None,
))
.await
.expect("router");
assert_error(unauthorized, StatusCode::UNAUTHORIZED, "unauthorized").await;
let view = app
.oneshot(request(
"GET",
format!("/api/sessions/{sid}"),
Some(&token),
Body::empty(),
None,
))
.await
.expect("router");
let view = response_json(view).await;
assert_eq!(view["participants"][0]["api_name"], "user_1");
assert!(!view.to_string().contains(&token));
}
#[tokio::test]
async fn endpoint_bootstrap_rejects_invalid_json_and_multipart_before_provider() {
let (state, sid, token, _) = seeded_state().await;
let app = build_router(state);
let unauth = app
.clone()
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/bootstrap/answer"),
None,
Body::from("{}"),
Some("application/json"),
))
.await
.expect("router");
assert_error(unauth, StatusCode::UNAUTHORIZED, "unauthorized").await;
let invalid_json = app
.clone()
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/bootstrap/answer"),
Some(&token),
Body::from("{}"),
Some("application/json"),
))
.await
.expect("router");
assert_error(invalid_json, StatusCode::BAD_REQUEST, "invalid_request").await;
let (type_header, payload) = multipart("audio/not-supported", b"x");
let invalid_audio = app
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/bootstrap/answer"),
Some(&token),
Body::from(payload),
Some(&type_header),
))
.await
.expect("router");
assert_error(invalid_audio, StatusCode::BAD_REQUEST, "invalid_request").await;
let (type_header, payload) = multipart("audio/webm", &[7; 33]);
let oversized = build_router(endpoint_state())
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/bootstrap/answer"),
Some(&token),
Body::from(payload),
Some(&type_header),
))
.await
.expect("router");
assert_error(oversized, StatusCode::NOT_FOUND, "session_not_found").await;
}
#[tokio::test]
async fn private_creator_audio_input_is_rejected_before_transcription() {
let state = endpoint_state();
let mut session = now_session("private".into());
let creator = add_participant(&mut session, "en".into(), true);
state
.sessions
.lock()
.await
.insert(session.id.clone(), Arc::new(Mutex::new(session)));
let (content_type, payload) = multipart("audio/webm", b"voice");
let response = build_router(state)
.oneshot(request(
"POST",
"/api/sessions/private/input/audio".into(),
Some(&creator.token),
Body::from(payload),
Some(&content_type),
))
.await
.expect("router");
assert_error(response, StatusCode::CONFLICT, "onboarding_required").await;
}
#[tokio::test]
async fn multipart_rejects_extra_and_duplicate_fields_and_bounds_total_body() {
let (state, sid, token, _) = seeded_state().await;
let app = build_router(state);
let boundary = "strict-multipart";
let audio = format!(
"--{boundary}\r\nContent-Disposition: form-data; name=\"audio\"; filename=\"v\"\r\nContent-Type: audio/webm\r\n\r\nx\r\n"
);
let extra =
format!("--{boundary}\r\nContent-Disposition: form-data; name=\"extra\"\r\n\r\nx\r\n");
for body in [
format!("{extra}{audio}--{boundary}--\r\n"),
format!("{audio}{extra}--{boundary}--\r\n"),
format!("{audio}{audio}--{boundary}--\r\n"),
format!(
"{audio}{extra}{}--{boundary}--\r\n",
"x".repeat(MAX_MULTIPART_OVERHEAD)
),
] {
let response = app
.clone()
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/input/audio"),
Some(&token),
Body::from(body),
Some(&format!("multipart/form-data; boundary={boundary}")),
))
.await
.expect("multipart response");
assert!(matches!(
response.status(),
StatusCode::BAD_REQUEST | StatusCode::PAYLOAD_TOO_LARGE
));
assert_eq!(response.headers()[header::CONTENT_TYPE], "application/json");
}
}
#[tokio::test]
async fn endpoint_inputs_are_accepted_without_floor_authority() {
let (state, sid, token, _) = seeded_state().await;
let session = state.sessions.lock().await[&sid].clone();
let second_token = {
let mut session = session.lock().await;
add_participant(&mut session, "en".into(), false).token
};
let app = build_router(state.clone());
let text = app
.clone()
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/input/text"),
Some(&token),
Body::from(r#"{"text":"ordered input"}"#),
Some("application/json"),
))
.await
.expect("router");
assert_eq!(text.status(), StatusCode::OK);
let second = app
.clone()
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/input/text"),
Some(&second_token),
Body::from(r#"{"text":"overlapping input"}"#),
Some("application/json"),
))
.await
.expect("router");
assert_eq!(second.status(), StatusCode::OK);
let session = state
.sessions
.lock()
.await
.get(&sid)
.expect("seeded session")
.clone();
let session = session.lock().await;
assert_eq!(
session
.pending_messages
.front()
.and_then(|m| m.content.as_ref().and_then(ChatContent::as_text)),
Some("Ada: ordered input")
);
assert_eq!(
session
.pending_messages
.get(1)
.and_then(|m| m.content.as_ref().and_then(ChatContent::as_text)),
Some("user_2: overlapping input")
);
}
#[tokio::test]
async fn endpoint_audio_delivery_completion_onboarding_and_leave() {
let (state, sid, token, _) = seeded_state().await;
let session = state
.sessions
.lock()
.await
.get(&sid)
.expect("seeded session")
.clone();
let (participant_id, delivery_id) = {
let mut s = session.lock().await;
let participant_id = s.participant_order[0].clone();
s.onboarding_audio = Some(OnboardingAudio {
id: "onboard".into(),
bytes: Some(vec![9]),
mime: "audio/ogg".into(),
text: "Name?".into(),
});
s.active_delivery = Some(Delivery {
id: "delivery".into(),
participant_id: participant_id.clone(),
participant_api_name: "user_1".into(),
text: "caption".into(),
audio: Some(vec![1, 2, 3]),
mime: "audio/mpeg".into(),
next_speaker: Some("user_1".into()),
ready: true,
});
s.delivery_acks
.insert("delivery".into(), HashSet::from([participant_id.clone()]));
(participant_id, "delivery".to_string())
};
let app = build_router(state.clone());
let onboarding = app
.clone()
.oneshot(request(
"GET",
format!("/api/sessions/{sid}/onboarding/audio"),
Some(&token),
Body::empty(),
None,
))
.await
.expect("router");
assert_error(onboarding, StatusCode::UNAUTHORIZED, "unauthorized").await;
let audio = app
.clone()
.oneshot(request(
"GET",
format!("/api/sessions/{sid}/deliveries/{delivery_id}/audio"),
Some(&token),
Body::empty(),
None,
))
.await
.expect("router");
assert_eq!(audio.headers()[header::CONTENT_TYPE], "audio/mpeg");
assert_eq!(
to_bytes(audio.into_body(), 10)
.await
.expect("audio")
.as_ref(),
[1, 2, 3]
);
let complete = app
.clone()
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/deliveries/{delivery_id}/complete"),
Some(&token),
Body::empty(),
None,
))
.await
.expect("router");
assert_eq!(complete.status(), StatusCode::OK);
let leave = app
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/leave"),
Some(&token),
Body::empty(),
None,
))
.await
.expect("router");
assert_eq!(leave.status(), StatusCode::OK);
let session = session.lock().await;
assert!(!session.participants[&participant_id].connected);
assert!(session.active_delivery.is_none());
}
#[tokio::test]
async fn endpoint_websocket_upgrade_requires_upgrade_headers_and_defers_token_to_hello() {
let (state, sid, token, _) = seeded_state().await;
let app = build_router(state);
let ordinary = app
.clone()
.oneshot(request(
"GET",
format!("/api/sessions/{sid}/ws"),
None,
Body::empty(),
None,
))
.await
.expect("router");
assert_eq!(ordinary.status(), StatusCode::BAD_REQUEST);
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("ephemeral listener");
let address = listener.local_addr().expect("listener address");
let server = tokio::spawn(async move {
axum::serve(listener, app)
.await
.expect("websocket test server");
});
let (mut socket, response) =
tokio_tungstenite::connect_async(format!("ws://{address}/api/sessions/{sid}/ws"))
.await
.expect("websocket upgrade");
assert_eq!(response.status(), StatusCode::SWITCHING_PROTOCOLS);
socket
.send(tokio_tungstenite::tungstenite::Message::Text(
format!(r#"{{"type":"hello","token":"{token}"}}"#).into(),
))
.await
.expect("hello");
let event = socket
.next()
.await
.expect("ready event")
.expect("websocket message")
.into_text()
.expect("event text");
assert_eq!(
serde_json::from_str::<Value>(&event).expect("event JSON")["type"],
"session.ready"
);
server.abort();
}
#[tokio::test]
async fn json_routes_reject_malformed_wrong_unknown_and_oversized_bodies() {
let (state, sid, token, _) = seeded_state().await;
let app = build_router(state);
let routes = [
("/api/sessions".to_string(), None),
(format!("/api/sessions/{sid}/join"), None),
(
format!("/api/sessions/{sid}/input/text"),
Some(token.as_str()),
),
];
for (path, credential) in routes {
for (body, content_type) in [
("{", Some("application/json")),
(r#"{"unknown":true}"#, Some("application/json")),
(r#"{"text":"x"}"#, Some("text/plain")),
(
&format!(r#"{{"text":"{}"}}"#, "x".repeat(MAX_JSON_BYTES)),
Some("application/json"),
),
] {
let response = app
.clone()
.oneshot(request(
"POST",
path.clone(),
credential,
Body::from(body.to_string()),
content_type,
))
.await
.expect("router response");
assert_error(response, StatusCode::BAD_REQUEST, "invalid_request").await;
}
}
}
#[test]
fn positive_number_and_voice_configuration_matrix_is_independent() {
assert_eq!(parse_positive("1", "LIMIT"), Ok(1));
assert!(parse_positive("0", "LIMIT").is_err());
assert!(parse_positive("no", "LIMIT").is_err());
for (stt, model, voice, valid) in [
("", "", "", true),
("stt", "", "", true),
("", "tts", "voice", true),
("stt", "tts", "voice", true),
("", "tts", "", false),
("", "", "voice", false),
] {
assert_eq!(
voice_config(stt.into(), model.into(), voice.into(), "mp3".into()).is_ok(),
valid
);
}
assert!(supported_tts_format("flac").is_err());
}
#[test]
fn connected_socket_prevents_expiration_but_socket_free_idle_session_expires() {
let mut session = ready();
session.sockets.clear();
session.last_activity_at = Instant::now() - Duration::from_secs(2);
assert!(session_expired(&session, Duration::from_secs(1)));
let participant_id = session.participant_order[0].clone();
let (tx, _rx) = mpsc::channel(1);
session.sockets.insert(
participant_id,
SocketConnection {
id: "live".into(),
tx,
},
);
assert!(!session_expired(&session, Duration::from_secs(1)));
}
#[tokio::test]
async fn bootstrap_text_limit_accepts_boundary_shape_and_rejects_oversize_before_provider() {
let state = endpoint_state();
let mut session = now_session("bootstrap-limit".into());
let creator = add_participant(&mut session, "en".into(), true);
let sid = session.id.clone();
state
.sessions
.lock()
.await
.insert(sid.clone(), Arc::new(Mutex::new(session)));
let app = build_router(state);
let oversized = app
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/bootstrap/answer"),
Some(&creator.token),
Body::from(format!(r#"{{"text":"{}"}}"#, "x".repeat(4001))),
Some("application/json"),
))
.await
.expect("oversized bootstrap response");
assert_error(oversized, StatusCode::BAD_REQUEST, "invalid_request").await;
assert_eq!(
limited_text("x".repeat(4000)).expect("boundary text").len(),
4000
);
assert!(limited_text("x".repeat(4001)).is_err());
}
#[test]
fn saturated_socket_is_removed_instead_of_losing_protocol_events_silently() {
let mut session = ready();
let participant_id = session.participant_order[0].clone();
let (tx, _rx) = mpsc::channel(1);
tx.try_send(ServerEvent::new("occupied", json!({})))
.expect("empty queue accepts the first event");
session.sockets.insert(
participant_id.clone(),
SocketConnection {
id: "connection".into(),
tx,
},
);
emit(
&mut session,
None,
ServerEvent::new("must-not-be-dropped", json!({})),
);
assert!(!session.sockets.contains_key(&participant_id));
assert!(!session.participants[&participant_id].connected);
}
#[test]
fn stale_model_turn_cannot_commit_after_newer_input_or_interrupt() {
let mut session = ready();
let snapshot = session.input_revision;
assert!(model_result_is_current(&session, snapshot));
session.input_revision += 1; // accepted input
assert!(!model_result_is_current(&session, snapshot));
let snapshot = session.input_revision;
session.input_revision += 1; // explicit interrupt
assert!(!model_result_is_current(&session, snapshot));
}
#[tokio::test]
async fn input_floor_assignment_emits_before_input_accepted_for_idle_text_and_audio() {
let mut session = ready();
let participant_id = session.participant_order[0].clone();
session.floor_holder = None;
let (tx, mut rx) = mpsc::channel(8);
session.sockets.insert(
participant_id.clone(),
SocketConnection {
id: "observed".into(),
tx,
},
);
set_floor(&mut session, &participant_id, false);
emit(
&mut session,
None,
ServerEvent::new("input.accepted", json!({"api_name":"user_1"})),
);
assert_eq!(rx.recv().await.expect("floor event").kind, "floor.changed");
assert_eq!(rx.recv().await.expect("input event").kind, "input.accepted");
session.floor_holder = None;
set_floor(&mut session, &participant_id, false);
assert_eq!(
rx.recv().await.expect("audio floor event").kind,
"floor.changed"
);
}
#[test]
fn saturation_of_last_delivery_recipient_returns_immediate_completion_signal() {
let mut session = ready();
let participant_id = session.participant_order[0].clone();
session.active_delivery = Some(delivery(participant_id.clone(), "active", "unused"));
session
.delivery_acks
.insert("active".into(), HashSet::from([participant_id.clone()]));
let (tx, _rx) = mpsc::channel(1);
tx.try_send(ServerEvent::new("occupied", json!({})))
.expect("fill socket queue");
session.sockets.insert(
participant_id,
SocketConnection {
id: "saturated".into(),
tx,
},
);
assert!(
!emit(
&mut session,
None,
ServerEvent::new("delivery.started", json!({}))
)
.start_delivery
);
assert!(session.active_delivery.is_none());
}
#[tokio::test]
async fn broadcast_acknowledgements_complete_once_and_advance_fifo() {
let (state, sid, token, _) = seeded_state().await;
let session = state.sessions.lock().await[&sid].clone();
let (second, third) = {
let mut s = session.lock().await;
let second = add_participant(&mut s, "en".into(), false);
let third = add_participant(&mut s, "en".into(), false);
for id in [&second.id, &third.id] {
s.participants.get_mut(id).expect("participant").onboarding = Onboarding::Complete;
test_connect(&mut s, id);
}
let recipients = HashSet::from([
s.participant_order[0].clone(),
second.id.clone(),
third.id.clone(),
]);
s.delivery_acks.insert("all".into(), recipients);
s.active_delivery = Some(Delivery {
id: "all".into(),
participant_id: "all".into(),
participant_api_name: "all".into(),
text: "one shared delivery".into(),
audio: None,
mime: "audio/mpeg".into(),
next_speaker: Some("user_2".into()),
ready: true,
});
s.deliveries
.push_back(delivery(second.id.clone(), "next", "unused"));
(second, third)
};
let app = build_router(state.clone());
for credential in [&token, &second.token] {
let response = app
.clone()
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/deliveries/all/complete"),
Some(credential),
Body::empty(),
None,
))
.await
.expect("ack");
assert_eq!(response.status(), StatusCode::OK);
assert_eq!(
session
.lock()
.await
.active_delivery
.as_ref()
.map(|d| d.id.as_str()),
Some("all")
);
}
let final_ack = app
.clone()
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/deliveries/all/complete"),
Some(&third.token),
Body::empty(),
None,
))
.await
.expect("final ack");
assert_eq!(final_ack.status(), StatusCode::OK);
start_delivery(state.clone(), sid.clone()).await;
let stale = app
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/deliveries/all/complete"),
Some(&third.token),
Body::empty(),
None,
))
.await
.expect("stale ack");
assert_error(stale, StatusCode::CONFLICT, "delivery_conflict").await;
let s = session.lock().await;
assert_eq!(s.floor_holder.as_deref(), Some(second.id.as_str()));
assert_eq!(
s.active_delivery.as_ref().map(|d| d.id.as_str()),
Some("next")
);
}
#[tokio::test]
async fn interrupt_clears_delivery_state_and_orders_protocol_events() {
let (state, sid, token, _) = seeded_state().await;
let session = state.sessions.lock().await[&sid].clone();
let (tx, mut rx) = mpsc::channel(8);
{
let mut s = session.lock().await;
let pid = s.participant_order[0].clone();
s.sockets.insert(
pid.clone(),
SocketConnection {
id: "socket".into(),
tx,
},
);
s.delivery_acks
.insert("active".into(), HashSet::from([pid.clone()]));
s.active_delivery = Some(delivery(pid.clone(), "active", "unused"));
s.deliveries.push_back(delivery(pid, "queued", "unused"));
}
let response = build_router(state)
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/interrupt"),
Some(&token),
Body::empty(),
None,
))
.await
.expect("interrupt");
assert_eq!(response.status(), StatusCode::OK);
let first = rx.recv().await.expect("interrupted");
let second = rx.recv().await.expect("floor");
assert_eq!(first.kind, "delivery.interrupted");
assert_eq!(second.kind, "floor.changed");
let s = session.lock().await;
assert!(
s.active_delivery.is_none() && s.deliveries.is_empty() && s.delivery_acks.is_empty()
);
}
#[tokio::test]
async fn broadcast_delivery_starts_with_one_id_for_every_recipient_and_authorizes_each() {
let (state, sid, token, _) = seeded_state().await;
let session = state.sessions.lock().await[&sid].clone();
let mut receivers = Vec::new();
let tokens = {
let mut s = session.lock().await;
for _ in 0..2 {
let participant = add_participant(&mut s, "en".into(), false);
s.participants
.get_mut(&participant.id)
.expect("participant")
.onboarding = Onboarding::Complete;
}
let ids = s.participant_order.clone();
for id in &ids {
let (tx, rx) = mpsc::channel(4);
s.sockets.insert(
id.clone(),
SocketConnection {
id: format!("socket-{id}"),
tx,
},
);
receivers.push(rx);
}
s.delivery_acks
.insert("shared".into(), ids.iter().cloned().collect());
s.deliveries.push_back(Delivery {
id: "shared".into(),
participant_id: "all".into(),
participant_api_name: "all".into(),
text: "shared".into(),
audio: Some(vec![1]),
mime: "audio/mpeg".into(),
next_speaker: None,
ready: true,
});
s.participants
.values()
.map(|p| p.token.clone())
.collect::<Vec<_>>()
};
start_delivery(state.clone(), sid.clone()).await;
for receiver in &mut receivers {
assert_eq!(
receiver.recv().await.expect("preparing").kind,
"delivery.preparing"
);
let event = receiver.recv().await.expect("started");
assert_eq!(event.kind, "delivery.started");
assert_eq!(event.data["delivery_id"], "shared");
}
let app = build_router(state);
for credential in tokens {
let response = app
.clone()
.oneshot(request(
"GET",
format!("/api/sessions/{sid}/deliveries/shared/audio"),
Some(&credential),
Body::empty(),
None,
))
.await
.expect("audio");
assert_eq!(response.status(), StatusCode::OK);
}
let response = app
.oneshot(request(
"GET",
format!("/api/sessions/{sid}/deliveries/shared/audio"),
Some("bad"),
Body::empty(),
None,
))
.await
.expect("unauthorized");
assert_error(response, StatusCode::UNAUTHORIZED, "unauthorized").await;
assert_ne!(token, "bad");
}
#[tokio::test]
async fn broadcast_disconnect_and_repeated_timeout_completion_are_safe() {
let (state, sid, _, _) = seeded_state().await;
let session = state.sessions.lock().await[&sid].clone();
let recipients = {
let mut s = session.lock().await;
let second = add_participant(&mut s, "en".into(), false);
s.participants
.get_mut(&second.id)
.expect("participant")
.onboarding = Onboarding::Complete;
let recipients: HashSet<_> = s.participant_order.iter().cloned().collect();
s.delivery_acks
.insert("disconnect".into(), recipients.clone());
s.active_delivery = Some(Delivery {
id: "disconnect".into(),
participant_id: "all".into(),
participant_api_name: "all".into(),
text: "x".into(),
audio: None,
mime: "audio/mpeg".into(),
next_speaker: None,
ready: true,
});
recipients
};
for recipient in recipients {
let outcome = {
let mut s = session.lock().await;
disconnect_participant(&mut s, &recipient, 512)
};
if let Some(id) = outcome.active_delivery {
complete_delivery(state.clone(), sid.clone(), id).await;
}
}
complete_delivery(state.clone(), sid.clone(), "disconnect".into()).await;
complete_delivery(state, sid, "disconnect".into()).await;
assert!(session.lock().await.active_delivery.is_none());
}
fn live_state() -> Arc<AppState> {
let config = Config::load()
.expect("live provider test requires configured provider environment variables");
Arc::new(AppState {
config,
http: reqwest::Client::builder()
.build()
.expect("HTTP client for live provider test"),
sessions: Mutex::new(HashMap::new()),
})
}
#[tokio::test]
#[ignore = "uses configured provider credentials and spends credits"]
async fn live_provider_forced_tool_chat_validates_against_ready_session() {
let state = live_state();
let session = ready();
let reply = chat(
&state,
session.messages.clone(),
tool_schema(),
true,
"live-provider-forced-tool",
)
.await
.expect("live chat provider request failed");
let (_, args) = parse_tool(reply)
.expect("live chat response must contain exactly one send_messages call");
assert!(
validate_tool(&session, args).is_ok(),
"live tool arguments must address the ready participant"
);
}
#[tokio::test]
#[ignore = "uses configured provider credentials and spends credits"]
async fn live_private_conversation_promotes_only_after_explicit_group_intent() {
let state = live_state();
let app = build_router(state.clone());
let created = app
.clone()
.oneshot(request(
"POST",
"/api/sessions".into(),
None,
Body::from(r#"{"browser_language":"en"}"#),
Some("application/json"),
))
.await
.expect("create provisional conversation");
assert_eq!(created.status(), StatusCode::OK);
let created = response_json(created).await;
let sid = created["session_id"].as_str().unwrap().to_string();
let token = created["participant_token"].as_str().unwrap().to_string();
let ordinary = app
.clone()
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/bootstrap/answer"),
Some(&token),
Body::from(r#"{"text":"Explain in one sentence why the sky appears blue."}"#),
Some("application/json"),
))
.await
.expect("ordinary private conversation turn");
assert_eq!(ordinary.status(), StatusCode::OK);
assert_eq!(response_json(ordinary).await["session_created"], false);
assert_eq!(
state.sessions.lock().await[&sid].lock().await.status,
SessionStatus::Onboarding
);
let promote = app
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/bootstrap/answer"),
Some(&token),
Body::from(r#"{"text":"Create a group chat for choosing a reliable launch plan. Call me Max and use English."}"#),
Some("application/json"),
))
.await
.expect("explicit group promotion turn");
assert_eq!(promote.status(), StatusCode::OK);
let promoted = response_json(promote).await;
assert_eq!(promoted["session_created"], true);
assert_eq!(promoted["session_id"], sid);
assert_eq!(
state.sessions.lock().await[&sid].lock().await.status,
SessionStatus::Active
);
}
#[tokio::test]
#[ignore = "uses configured provider credentials and spends credits"]
async fn live_provider_three_participant_tool_argument_validation() {
let state = live_state();
let mut session = ready();
for name in ["Grace", "Linus"] {
let participant = add_participant(&mut session, "en".into(), false);
let participant = session.participants.get_mut(&participant.id).unwrap();
participant.display_name = Some(name.into());
participant.onboarding = Onboarding::Complete;
}
let system_prompt = "You facilitate a three-person conversation. Every response must use send_messages to address user_1, user_2, and user_3, either separately or with one dispatch to all. Keep each response concise.".to_string();
session.messages[0].content = Some(ChatContent::Text(system_prompt));
let participant_ids = session.participant_order.clone();
for round in 1..=3 {
for (index, participant_id) in participant_ids.iter().enumerate() {
let participant = &session.participants[participant_id];
session.messages.push(ChatMessage {
role: "user".into(),
name: Some(participant.api_name.clone()),
content: Some(ChatContent::Text(format!(
"{}: round {round}, participant {} contributes option {}.",
participant.display_name.as_deref().unwrap(),
index + 1,
round * 10 + index
))),
tool_calls: None,
tool_call_id: None,
});
}
let reply = chat(
&state,
session.messages.clone(),
tool_schema(),
true,
"live-three-participant-dispatch",
)
.await
.expect("live three-participant chat request");
let (call, args) = parse_tool(reply).expect("valid send_messages tool call");
let dispatches =
validate_tool(&session, args).expect("valid three-participant dispatches");
assert!(
dispatches.iter().any(|dispatch| dispatch.reply_to == "all")
|| ["user_1", "user_2", "user_3"].iter().all(|recipient| {
dispatches
.iter()
.any(|dispatch| dispatch.reply_to == *recipient)
})
);
let call_id = call.id.clone();
session.messages.push(ChatMessage {
role: "assistant".into(),
name: None,
content: None,
tool_calls: Some(vec![call]),
tool_call_id: None,
});
session.messages.push(ChatMessage {
role: "tool".into(),
name: None,
content: Some(ChatContent::Text(
json!({"status":"success","queued":dispatches.len()}).to_string(),
)),
tool_calls: None,
tool_call_id: Some(call_id),
});
}
}
#[tokio::test]
#[ignore = "uses configured provider credentials and spends credits"]
async fn live_real_three_participant_server_flow() {
let configured = live_state();
let mut config = configured.config.clone();
config.batch = Duration::from_secs(3600);
config.tts = None;
let state = Arc::new(AppState {
config,
http: configured.http.clone(),
sessions: Mutex::new(HashMap::new()),
});
let app = build_router(state.clone());
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("ephemeral listener");
let address = listener.local_addr().expect("listener address");
let server_app = app.clone();
let server = tokio::spawn(async move {
axum::serve(listener, server_app)
.await
.expect("test server");
});
let created = app
.clone()
.oneshot(request(
"POST",
"/api/sessions".into(),
None,
Body::from(r#"{"browser_language":"en"}"#),
Some("application/json"),
))
.await
.expect("create");
assert_eq!(created.status(), StatusCode::OK);
let created = response_json(created).await;
let sid = created["session_id"]
.as_str()
.expect("session id")
.to_string();
let creator = created["participant_token"]
.as_str()
.expect("creator token")
.to_string();
let ordinary = app.clone().oneshot(request("POST", format!("/api/sessions/{sid}/bootstrap/answer"), Some(&creator), Body::from(r#"{"text":"I still don't understand how to choose between two options; please help."}"#), Some("application/json"))).await.expect("ordinary help");
assert_eq!(ordinary.status(), StatusCode::OK);
assert!(
!response_json(ordinary).await["session_created"]
.as_bool()
.expect("boolean")
);
let promoted = app.clone().oneshot(request("POST", format!("/api/sessions/{sid}/bootstrap/answer"), Some(&creator), Body::from(r#"{"text":"Create a general group game for three people: collaboratively choose clues, take turns proposing answers, keep score fairly, and explain the rules clearly."}"#), Some("application/json"))).await.expect("promotion");
assert_eq!(promoted.status(), StatusCode::OK);
assert!(
response_json(promoted).await["session_created"]
.as_bool()
.expect("boolean")
);
let session = state.sessions.lock().await[&sid].clone();
{
let s = session.lock().await;
let prompt = s
.messages
.first()
.and_then(|message| message.content.as_ref())
.and_then(ChatContent::as_text)
.expect("group system prompt");
assert!(prompt.len() > DEFAULT_GROUP_SESSION_PROTOCOL_PROMPT.len());
assert!(prompt.contains(&state.config.group_session_protocol_prompt));
}
let mut joined = Vec::new();
for language in ["en", "en"] {
let response = app
.clone()
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/join"),
None,
Body::from(format!(r#"{{"browser_language":"{language}"}}"#)),
Some("application/json"),
))
.await
.expect("join");
assert_eq!(response.status(), StatusCode::OK);
joined.push(
response_json(response).await["participant_token"]
.as_str()
.expect("join token")
.to_string(),
);
}
let mut sockets = Vec::new();
for credential in std::iter::once(&creator).chain(joined.iter()) {
let (mut socket, _) = connect_async(format!("ws://{address}/api/sessions/{sid}/ws"))
.await
.expect("participant WebSocket");
socket
.send(TungsteniteMessage::Text(
json!({"type":"hello","token":credential})
.to_string()
.into(),
))
.await
.expect("hello");
let ready = tokio::time::timeout(Duration::from_secs(5), socket.next())
.await
.expect("session.ready timeout")
.expect("WebSocket open")
.expect("WebSocket event");
assert_eq!(
serde_json::from_str::<Value>(ready.into_text().expect("text event").as_ref())
.expect("event JSON")["type"],
"session.ready"
);
sockets.push(socket);
}
let mut interrupted_delivery = false;
for (credential, text) in [
(&joined[0], "Hello, I have not chosen a display name yet."),
(&joined[0], "My name is River; I propose a riddle."),
(&joined[1], "I'm Sol. I am ready to play."),
(&creator, "Let's begin now."),
(&creator, "Pause the game and let me clarify the rules."),
] {
let response = app
.clone()
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/input/text"),
Some(credential),
Body::from(format!(
r#"{{"text":{}}}"#,
serde_json::to_string(text).expect("text JSON")
)),
Some("application/json"),
))
.await
.expect("input");
assert_eq!(response.status(), StatusCode::OK);
run_model(state.clone(), sid.clone()).await;
start_delivery(state.clone(), sid.clone()).await;
let acknowledgements = {
let s = session.lock().await;
s.active_delivery.as_ref().map(|delivery| {
(
delivery.id.clone(),
s.delivery_acks
.get(&delivery.id)
.cloned()
.unwrap_or_default(),
)
})
};
if let Some((delivery_id, recipients)) = acknowledgements {
let recipient_indices = {
let s = session.lock().await;
recipients
.iter()
.map(|recipient| {
s.participant_order
.iter()
.position(|id| id == recipient)
.expect("delivery recipient is a participant")
})
.collect::<Vec<_>>()
};
for index in &recipient_indices {
let event = tokio::time::timeout(Duration::from_secs(15), async {
loop {
let message = sockets[*index]
.next()
.await
.expect("WebSocket remains open")
.expect("WebSocket event");
let event: Value = serde_json::from_str(
message.into_text().expect("text event").as_ref(),
)
.expect("event JSON");
if event["type"] == "delivery.started" {
return event;
}
}
})
.await
.expect("delivery.started timeout");
assert_eq!(event["delivery_id"], delivery_id);
}
for index in (0..sockets.len()).filter(|index| !recipient_indices.contains(index)) {
let unexpected = tokio::time::timeout(Duration::from_millis(250), async {
loop {
let message = sockets[index]
.next()
.await
.expect("WebSocket remains open")
.expect("WebSocket event");
let event: Value = serde_json::from_str(
message.into_text().expect("text event").as_ref(),
)
.expect("event JSON");
if event["type"] == "delivery.started" {
return event;
}
}
})
.await;
assert!(
unexpected.is_err(),
"targeted delivery reached a non-recipient"
);
}
if !interrupted_delivery {
let response = app
.clone()
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/interrupt"),
Some(&creator),
Body::empty(),
None,
))
.await
.expect("interrupt response");
assert_eq!(response.status(), StatusCode::OK);
for socket in &mut sockets {
let interrupted = tokio::time::timeout(Duration::from_secs(5), async {
loop {
let message = socket
.next()
.await
.expect("WebSocket remains open")
.expect("WebSocket event");
let event: Value = serde_json::from_str(
message.into_text().expect("text event").as_ref(),
)
.expect("event JSON");
if event["type"] == "delivery.interrupted" {
return event;
}
}
})
.await
.expect("delivery.interrupted timeout");
assert_eq!(interrupted["delivery_id"], delivery_id);
let floor = tokio::time::timeout(Duration::from_secs(5), socket.next())
.await
.expect("floor.changed timeout")
.expect("WebSocket remains open")
.expect("WebSocket event");
assert_eq!(
serde_json::from_str::<Value>(
floor.into_text().expect("text event").as_ref()
)
.expect("event JSON")["type"],
"floor.changed"
);
}
let stale = app
.clone()
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/deliveries/{delivery_id}/complete"),
Some(&creator),
Body::empty(),
None,
))
.await
.expect("stale completion response");
assert_eq!(stale.status(), StatusCode::CONFLICT);
interrupted_delivery = true;
continue;
}
let credentials = {
let s = session.lock().await;
recipients
.iter()
.map(|recipient| {
let participant = s.participants.get(recipient).expect("recipient");
assert!(participant.api_name.starts_with("user_"));
participant.token.clone()
})
.collect::<Vec<_>>()
};
for credential in credentials {
let response = app
.clone()
.oneshot(request(
"POST",
format!("/api/sessions/{sid}/deliveries/{delivery_id}/complete"),
Some(&credential),
Body::empty(),
None,
))
.await
.expect("completion");
assert_eq!(response.status(), StatusCode::OK);
}
assert_ne!(
session
.lock()
.await
.active_delivery
.as_ref()
.map(|delivery| delivery.id.as_str()),
Some(delivery_id.as_str())
);
}
}
let s = session.lock().await;
assert!(
s.participants.len() == 3 && s.messages.iter().any(|message| message.role == "tool")
);
drop(s);
server.abort();
}
#[tokio::test]
#[ignore = "uses configured provider credentials and spends credits"]
async fn live_direct_chat_audio_without_stt() {
let configured = live_state();
let mut config = configured.config.clone();
config.stt_model = None;
config.tts = None;
config.batch = Duration::from_secs(3600);
let state = Arc::new(AppState {
config,
http: configured.http.clone(),
sessions: Mutex::new(HashMap::new()),
});
let mut session = ready();
session.id = "live-direct-audio".into();
let sid = session.id.clone();
let pid = session.participant_order[0].clone();
let token = session.participants[&pid].token.clone();
state
.sessions
.lock()
.await
.insert(sid.clone(), Arc::new(Mutex::new(session)));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0")
.await
.expect("ephemeral listener");
let address = listener.local_addr().expect("listener address");
let server_state = state.clone();
let server = tokio::spawn(async move {
axum::serve(listener, build_router(server_state))
.await
.expect("test server");
});
let base = format!("http://{address}");
let (mut socket, _) = connect_async(format!("ws://{address}/api/sessions/{sid}/ws"))
.await
.expect("participant WebSocket");
socket
.send(TungsteniteMessage::Text(
json!({"type":"hello","token":token}).to_string().into(),
))
.await
.expect("hello");
let ready = tokio::time::timeout(Duration::from_secs(5), socket.next())
.await
.expect("session.ready timeout")
.expect("WebSocket open")
.expect("WebSocket event");
assert_eq!(
serde_json::from_str::<Value>(ready.into_text().expect("text event").as_ref())
.expect("event JSON")["type"],
"session.ready"
);
let wav = {
let sample_rate = 8_000_u32;
let samples = sample_rate / 20;
let data_len = samples * 2;
let mut bytes = b"RIFF".to_vec();
bytes.extend_from_slice(&(36 + data_len).to_le_bytes());
bytes.extend_from_slice(b"WAVEfmt ");
bytes.extend_from_slice(&16_u32.to_le_bytes());
bytes.extend_from_slice(&1_u16.to_le_bytes());
bytes.extend_from_slice(&1_u16.to_le_bytes());
bytes.extend_from_slice(&sample_rate.to_le_bytes());
bytes.extend_from_slice(&(sample_rate * 2).to_le_bytes());
bytes.extend_from_slice(&2_u16.to_le_bytes());
bytes.extend_from_slice(&16_u16.to_le_bytes());
bytes.extend_from_slice(b"data");
bytes.extend_from_slice(&data_len.to_le_bytes());
for sample in 0..samples {
let value = if sample % 16 < 8 {
4_000_i16
} else {
-4_000_i16
};
bytes.extend_from_slice(&value.to_le_bytes());
}
bytes
};
let boundary = "live-audio-boundary";
let mut body = format!("--{boundary}\r\nContent-Disposition: form-data; name=\"audio\"; filename=\"voice.wav\"\r\nContent-Type: audio/wav\r\n\r\n").into_bytes();
body.extend_from_slice(&wav);
body.extend_from_slice(format!("\r\n--{boundary}--\r\n").as_bytes());
let response = reqwest::Client::new()
.post(format!("{base}/api/sessions/{sid}/input/audio"))
.bearer_auth(&token)
.header(
"content-type",
format!("multipart/form-data; boundary={boundary}"),
)
.body(body)
.send()
.await
.expect("direct audio route");
assert_eq!(response.status(), StatusCode::OK);
run_model(state.clone(), sid.clone()).await;
let event = tokio::time::timeout(Duration::from_secs(15), async {
loop {
let message = socket
.next()
.await
.expect("WebSocket remains open")
.expect("WebSocket event");
let event: Value =
serde_json::from_str(message.into_text().expect("text event").as_ref())
.expect("event JSON");
if event["type"] == "delivery.started" {
return event;
}
}
})
.await
.expect("delivery.started timeout");
assert!(event["delivery_id"].is_string());
server.abort();
}
#[tokio::test]
#[ignore = "uses configured provider credentials and spends credits"]
async fn live_provider_tts_then_stt_returns_nonempty_transcript() {
let state = live_state();
if state.config.tts.is_none() || state.config.stt_model.is_none() {
eprintln!("skipping optional STT/TTS round trip: both settings are required");
return;
}
let (audio, mime) = speak(&state, "Please confirm receipt.")
.await
.expect("live TTS provider request failed");
assert!(!audio.is_empty());
let format = AudioFormat::from_mime(&mime)
.map(|format| format.as_str())
.or_else(|| {
match state
.config
.tts
.as_ref()
.expect("checked above")
.tts_format
.as_str()
{
"mp3" => Some("mp3"),
"wav" => Some("wav"),
"ogg" => Some("ogg"),
"webm" => Some("webm"),
"m4a" => Some("m4a"),
_ => None,
}
})
.expect("TTS_FORMAT must map to a supported transcription format");
let transcript = transcribe(&state, audio, format, &state.config.default_language)
.await
.expect("live STT provider request failed");
assert!(!transcript.trim().is_empty());
}
}