use std::collections::HashMap;
use std::sync::Arc;
use std::time::{Duration, Instant};
use crate::host::{LiveAdapterHost, LiveChannelId, ObservationOutcome};
use axum::Router;
use axum::extract::ws::{CloseFrame, Message as WsMessage, WebSocket, close_code};
use axum::extract::{Query, State, WebSocketUpgrade};
use axum::response::IntoResponse;
use axum::routing::get;
use meerkat_contracts::WireLiveAdapterObservation;
use meerkat_core::live_adapter::LiveInputChunk;
use tokio::sync::Mutex;
use uuid::Uuid;
pub const LIVE_WS_PATH: &str = "/live/ws";
pub const TOKEN_TTL: Duration = Duration::from_secs(60);
#[derive(Debug, serde::Serialize)]
struct WsErrorFrame {
error: String,
#[serde(skip_serializing_if = "Option::is_none")]
reason: Option<String>,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq)]
pub enum BinaryFormat {
Pcm24kMono,
}
impl BinaryFormat {
fn parse(s: &str) -> Option<Self> {
match s {
"pcm_24k_mono" => Some(Self::Pcm24kMono),
_ => None,
}
}
fn sample_rate_hz(self) -> u32 {
match self {
Self::Pcm24kMono => 24_000,
}
}
fn channels(self) -> u16 {
match self {
Self::Pcm24kMono => 1,
}
}
}
#[derive(Debug, Clone, PartialEq, Eq, Hash)]
pub struct LiveTokenString(String);
impl LiveTokenString {
pub fn new(s: impl Into<String>) -> Result<Self, TokenParseError> {
let s = s.into();
if s.is_empty() {
return Err(TokenParseError::Empty);
}
for (idx, b) in s.bytes().enumerate() {
let ok = b.is_ascii_alphanumeric() || matches!(b, b'-' | b'_' | b'.' | b'~');
if !ok {
return Err(TokenParseError::InvalidByte { idx, byte: b });
}
}
Ok(Self(s))
}
pub(crate) fn random() -> Self {
Self(Uuid::new_v4().to_string())
}
pub fn as_str(&self) -> &str {
&self.0
}
}
impl std::fmt::Display for LiveTokenString {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
f.write_str(&self.0)
}
}
#[derive(Debug, thiserror::Error)]
pub enum TokenParseError {
#[error("token is empty")]
Empty,
#[error("token contains non-URL-safe byte 0x{byte:02x} at index {idx}")]
InvalidByte { idx: usize, byte: u8 },
}
#[derive(Debug, thiserror::Error, PartialEq, Eq)]
pub enum TokenConsumeError {
#[error("token not found")]
NotFound,
#[error("token expired")]
Expired,
#[error("token bound to a different channel")]
ChannelMismatch,
}
struct PendingToken {
channel_id: LiveChannelId,
expires_at: Instant,
}
pub struct LiveWsState {
host: Arc<LiveAdapterHost>,
pending_tokens: Mutex<HashMap<LiveTokenString, PendingToken>>,
token_ttl: Duration,
}
impl LiveWsState {
pub fn new(host: Arc<LiveAdapterHost>) -> Self {
Self::with_token_ttl(host, TOKEN_TTL)
}
pub fn with_token_ttl(host: Arc<LiveAdapterHost>, token_ttl: Duration) -> Self {
Self {
host,
pending_tokens: Mutex::new(HashMap::new()),
token_ttl,
}
}
pub fn host(&self) -> &Arc<LiveAdapterHost> {
&self.host
}
pub async fn mint_token(&self, channel_id: LiveChannelId) -> LiveTokenString {
let token = LiveTokenString::random();
let expires_at = Instant::now() + self.token_ttl;
let mut guard = self.pending_tokens.lock().await;
reap_expired(&mut guard);
guard.insert(
token.clone(),
PendingToken {
channel_id,
expires_at,
},
);
token
}
pub async fn consume_token(
&self,
token: &str,
expected_channel: &LiveChannelId,
) -> Result<LiveChannelId, TokenConsumeError> {
let key = LiveTokenString::new(token).map_err(|_| TokenConsumeError::NotFound)?;
let mut guard = self.pending_tokens.lock().await;
reap_expired(&mut guard);
match guard.remove(&key) {
Some(pending) => {
if pending.expires_at <= Instant::now() {
Err(TokenConsumeError::Expired)
} else if &pending.channel_id != expected_channel {
Err(TokenConsumeError::ChannelMismatch)
} else {
Ok(pending.channel_id)
}
}
None => Err(TokenConsumeError::NotFound),
}
}
#[cfg(test)]
async fn pending_token_count(&self) -> usize {
self.pending_tokens.lock().await.len()
}
}
fn reap_expired(map: &mut HashMap<LiveTokenString, PendingToken>) {
let now = Instant::now();
map.retain(|_, p| p.expires_at > now);
}
#[derive(serde::Deserialize)]
pub struct WsConnectParams {
pub token: String,
pub channel: String,
#[serde(default)]
pub format: Option<String>,
}
pub fn live_ws_router(state: Arc<LiveWsState>) -> Router {
Router::new()
.route(LIVE_WS_PATH, get(ws_upgrade))
.with_state(state)
}
async fn ws_upgrade(
ws: WebSocketUpgrade,
Query(params): Query<WsConnectParams>,
State(state): State<Arc<LiveWsState>>,
) -> impl IntoResponse {
let WsConnectParams {
token,
channel,
format,
} = params;
let binary_format = format.as_deref().and_then(BinaryFormat::parse);
let expected_channel = LiveChannelId::new(channel);
ws.on_upgrade(move |socket| {
handle_live_socket(socket, token, expected_channel, binary_format, state)
})
}
async fn close_with(socket: &mut WebSocket, code: u16, reason: &str) {
let _ = socket
.send(WsMessage::Close(Some(CloseFrame {
code,
reason: reason.to_owned().into(),
})))
.await;
}
fn live_adapter_error_code_slug(code: &meerkat_core::live_adapter::LiveAdapterErrorCode) -> String {
serde_json::to_value(code)
.ok()
.and_then(|v| v.get("code").and_then(|c| c.as_str()).map(str::to_owned))
.unwrap_or_else(|| "unknown".to_owned())
}
async fn handle_live_socket(
mut socket: WebSocket,
token: String,
expected_channel: LiveChannelId,
binary_format: Option<BinaryFormat>,
state: Arc<LiveWsState>,
) {
let channel_id = match state.consume_token(&token, &expected_channel).await {
Ok(id) => id,
Err(err) => {
let err_json = serde_json::to_string(&WsErrorFrame {
error: "invalid_token".into(),
reason: Some(err.to_string()),
})
.unwrap_or_default();
let _ = socket.send(WsMessage::Text(err_json.into())).await;
close_with(&mut socket, close_code::POLICY, "invalid_token").await;
return;
}
};
tracing::info!(channel = %channel_id, "live WebSocket connected");
let mut observation_fut = Box::pin(state.host.next_observation_raw(&channel_id));
loop {
tokio::select! {
client_msg = socket.recv() => {
match client_msg {
Some(Ok(WsMessage::Text(text))) => {
match serde_json::from_str::<LiveInputChunk>(text.as_str()) {
Ok(chunk) => {
if let Err(err) = state.host.send_input(&channel_id, chunk).await {
tracing::warn!(channel = %channel_id, error = %err, "send_input failed");
let err_json = serde_json::to_string(&WsErrorFrame {
error: err.to_string(),
reason: None,
}).unwrap_or_default();
let _ = socket.send(WsMessage::Text(err_json.into())).await;
}
}
Err(parse_err) => {
tracing::warn!(
channel = %channel_id,
error = %parse_err,
"invalid WS text frame; closing"
);
close_with(&mut socket, close_code::INVALID, "invalid_frame").await;
break;
}
}
}
Some(Ok(WsMessage::Binary(data))) => {
let Some(fmt) = binary_format else {
tracing::warn!(
channel = %channel_id,
"binary frame received before format negotiation; closing"
);
close_with(
&mut socket,
close_code::POLICY,
"binary_format_unnegotiated",
).await;
break;
};
let chunk = LiveInputChunk::Audio {
data: data.to_vec(),
sample_rate_hz: fmt.sample_rate_hz(),
channels: fmt.channels(),
};
if let Err(err) = state.host.send_input(&channel_id, chunk).await {
tracing::warn!(channel = %channel_id, error = %err, "binary send_input failed");
let err_json = serde_json::to_string(&WsErrorFrame {
error: err.to_string(),
reason: None,
}).unwrap_or_default();
let _ = socket.send(WsMessage::Text(err_json.into())).await;
}
}
Some(Ok(WsMessage::Close(_))) | None => break,
Some(Ok(WsMessage::Ping(_))) => {
}
Some(Ok(other)) => {
tracing::warn!(
channel = %channel_id,
kind = ?std::mem::discriminant(&other),
"unsupported WS frame; closing"
);
close_with(&mut socket, close_code::UNSUPPORTED, "unsupported_frame").await;
break;
}
Some(Err(err)) => {
tracing::warn!(channel = %channel_id, error = %err, "WS recv error");
break;
}
}
}
observation = &mut observation_fut => {
observation_fut = Box::pin(state.host.next_observation_raw(&channel_id));
match observation {
Ok(Some(obs)) => {
let wire_obs = WireLiveAdapterObservation::from(obs.clone());
let send_ok = match serde_json::to_string(&wire_obs) {
Ok(json) => socket.send(WsMessage::Text(json.into())).await.is_ok(),
Err(_) => true,
};
let outcome = state.host.apply_observation(&channel_id, &obs).await;
if !send_ok {
break;
}
match outcome {
Ok(ObservationOutcome::Terminal { code }) => {
let slug = live_adapter_error_code_slug(&code);
let reason = format!("terminal:{slug}");
close_with(&mut socket, close_code::POLICY, &reason).await;
break;
}
Ok(ObservationOutcome::CommandRejected { code, message }) => {
tracing::info!(
channel = %channel_id,
?code,
%message,
"live command rejected; channel remains open"
);
}
Ok(_) => {}
Err(err) => {
tracing::warn!(
channel = %channel_id,
error = %err,
"apply_observation failed; closing channel"
);
break;
}
}
}
Ok(None) => break,
Err(_) => break,
}
}
}
}
tracing::info!(channel = %channel_id, "live WebSocket disconnected");
let _ = state.host.close_channel(&channel_id).await;
}
pub async fn serve_live_ws_listener(
listener: tokio::net::TcpListener,
state: Arc<LiveWsState>,
) -> Result<(), std::io::Error> {
let app = live_ws_router(state);
axum::serve(listener, app).await
}
#[cfg(test)]
#[allow(clippy::unwrap_used, clippy::expect_used, clippy::panic)]
mod tests {
use super::*;
use crate::host::NoOpProjectionSink;
#[test]
fn token_string_accepts_url_safe_alphabet() {
LiveTokenString::new("abcXYZ-_.~0123").unwrap();
LiveTokenString::new("550e8400-e29b-41d4-a716-446655440000").unwrap();
}
#[test]
fn token_string_rejects_unsafe_bytes() {
assert!(matches!(
LiveTokenString::new("ab cd"),
Err(TokenParseError::InvalidByte { .. })
));
assert!(matches!(
LiveTokenString::new("ab/cd"),
Err(TokenParseError::InvalidByte { .. })
));
assert!(matches!(
LiveTokenString::new("ab%20cd"),
Err(TokenParseError::InvalidByte { .. })
));
assert!(matches!(
LiveTokenString::new(""),
Err(TokenParseError::Empty)
));
}
#[tokio::test]
async fn mint_and_consume_token() {
let host = Arc::new(LiveAdapterHost::new(Arc::new(NoOpProjectionSink)));
let state = LiveWsState::new(host);
let channel_id = LiveChannelId::new("test_ch");
let token = state.mint_token(channel_id.clone()).await;
assert!(!token.as_str().is_empty());
let consumed = state.consume_token(token.as_str(), &channel_id).await;
assert_eq!(consumed.unwrap(), channel_id);
assert_eq!(
state
.consume_token(token.as_str(), &channel_id)
.await
.unwrap_err(),
TokenConsumeError::NotFound,
);
}
#[tokio::test]
async fn consume_unknown_token_returns_not_found() {
let host = Arc::new(LiveAdapterHost::new(Arc::new(NoOpProjectionSink)));
let state = LiveWsState::new(host);
let any_channel = LiveChannelId::new("any");
assert_eq!(
state
.consume_token("bogus", &any_channel)
.await
.unwrap_err(),
TokenConsumeError::NotFound,
);
}
#[tokio::test]
async fn consume_malformed_token_returns_not_found() {
let host = Arc::new(LiveAdapterHost::new(Arc::new(NoOpProjectionSink)));
let state = LiveWsState::new(host);
let any_channel = LiveChannelId::new("any");
assert_eq!(
state
.consume_token("has spaces", &any_channel)
.await
.unwrap_err(),
TokenConsumeError::NotFound,
);
}
#[tokio::test]
async fn consume_token_with_wrong_channel_rejects() {
let host = Arc::new(LiveAdapterHost::new(Arc::new(NoOpProjectionSink)));
let state = LiveWsState::new(host);
let channel_a = LiveChannelId::new("ch_a");
let channel_b = LiveChannelId::new("ch_b");
let token = state.mint_token(channel_a.clone()).await;
assert_eq!(
state
.consume_token(token.as_str(), &channel_b)
.await
.unwrap_err(),
TokenConsumeError::ChannelMismatch,
);
assert_eq!(
state
.consume_token(token.as_str(), &channel_a)
.await
.unwrap_err(),
TokenConsumeError::NotFound,
"token must remain consumed after ChannelMismatch (no retry)"
);
}
#[tokio::test]
async fn token_expires_after_ttl_and_is_reaped() {
let host = Arc::new(LiveAdapterHost::new(Arc::new(NoOpProjectionSink)));
let state = LiveWsState::with_token_ttl(host, Duration::from_millis(50));
let channel_id = LiveChannelId::new("ttl_ch");
let token = state.mint_token(channel_id.clone()).await;
assert_eq!(state.pending_token_count().await, 1);
tokio::time::sleep(Duration::from_millis(120)).await;
let err = state
.consume_token(token.as_str(), &channel_id)
.await
.unwrap_err();
assert!(
matches!(
err,
TokenConsumeError::NotFound | TokenConsumeError::Expired
),
"unexpected error: {err:?}"
);
assert_eq!(state.pending_token_count().await, 0);
}
#[tokio::test]
async fn unrelated_mint_reaps_expired_tokens() {
let host = Arc::new(LiveAdapterHost::new(Arc::new(NoOpProjectionSink)));
let state = LiveWsState::with_token_ttl(host, Duration::from_millis(40));
let _stale = state.mint_token(LiveChannelId::new("stale")).await;
assert_eq!(state.pending_token_count().await, 1);
tokio::time::sleep(Duration::from_millis(80)).await;
let _fresh = state.mint_token(LiveChannelId::new("fresh")).await;
assert_eq!(state.pending_token_count().await, 1);
}
#[tokio::test]
async fn websocket_roundtrip_with_token() {
let host = Arc::new(LiveAdapterHost::new(Arc::new(NoOpProjectionSink)));
let session_id = meerkat_core::types::SessionId::new();
let channel_id = host.open_channel(session_id).await.unwrap();
host.attach_adapter(&channel_id, Arc::new(IdleAdapter))
.await
.unwrap();
host.apply_status_update(
&channel_id,
meerkat_core::live_adapter::LiveAdapterStatus::Ready,
)
.await
.unwrap();
let state = Arc::new(LiveWsState::new(host));
let token = state.mint_token(channel_id.clone()).await;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let ws_state = Arc::clone(&state);
let server_handle =
tokio::spawn(async move { serve_live_ws_listener(listener, ws_state).await });
let url = format!("ws://{addr}{LIVE_WS_PATH}?token={token}&channel={channel_id}");
let (ws_stream, _) = tokio_tungstenite::connect_async(&url).await.unwrap();
let (mut write, _read) = futures::StreamExt::split(ws_stream);
use futures::SinkExt;
use tokio_tungstenite::tungstenite::Message;
let input = serde_json::json!({"kind": "text", "text": "hello"});
write
.send(Message::Text(input.to_string().into()))
.await
.unwrap();
write.send(Message::Close(None)).await.unwrap();
server_handle.abort();
}
#[test]
fn live_adapter_error_code_slug_emits_serde_tag() {
use meerkat_core::live_adapter::LiveAdapterErrorCode;
assert_eq!(
live_adapter_error_code_slug(&LiveAdapterErrorCode::ProviderError),
"provider_error"
);
assert_eq!(
live_adapter_error_code_slug(&LiveAdapterErrorCode::ConnectionLost),
"connection_lost"
);
assert_eq!(
live_adapter_error_code_slug(&LiveAdapterErrorCode::Other {
raw: "custom".into()
}),
"other"
);
assert_eq!(
live_adapter_error_code_slug(&LiveAdapterErrorCode::ConfigRejected {
reason: meerkat_core::live_adapter::LiveConfigRejectionReason::Other {
detail: "model swap requires close + reopen".into(),
},
}),
"config_rejected"
);
}
struct IdleAdapter;
#[async_trait::async_trait]
impl meerkat_core::live_adapter::LiveAdapter for IdleAdapter {
async fn send_command(
&self,
_command: meerkat_core::live_adapter::LiveAdapterCommand,
) -> Result<(), meerkat_core::live_adapter::LiveAdapterError> {
Ok(())
}
async fn next_observation(
&self,
) -> Result<
Option<meerkat_core::live_adapter::LiveAdapterObservation>,
meerkat_core::live_adapter::LiveAdapterError,
> {
std::future::pending().await
}
fn status(&self) -> meerkat_core::live_adapter::LiveAdapterStatus {
meerkat_core::live_adapter::LiveAdapterStatus::Ready
}
async fn close(&self) -> Result<(), meerkat_core::live_adapter::LiveAdapterError> {
Ok(())
}
}
struct ScriptedAdapter {
observation: tokio::sync::Mutex<Option<meerkat_core::live_adapter::LiveAdapterObservation>>,
}
impl ScriptedAdapter {
fn new(observation: meerkat_core::live_adapter::LiveAdapterObservation) -> Self {
Self {
observation: tokio::sync::Mutex::new(Some(observation)),
}
}
}
#[async_trait::async_trait]
impl meerkat_core::live_adapter::LiveAdapter for ScriptedAdapter {
async fn send_command(
&self,
_command: meerkat_core::live_adapter::LiveAdapterCommand,
) -> Result<(), meerkat_core::live_adapter::LiveAdapterError> {
Ok(())
}
async fn next_observation(
&self,
) -> Result<
Option<meerkat_core::live_adapter::LiveAdapterObservation>,
meerkat_core::live_adapter::LiveAdapterError,
> {
let mut slot = self.observation.lock().await;
if let Some(obs) = slot.take() {
Ok(Some(obs))
} else {
std::future::pending().await
}
}
fn status(&self) -> meerkat_core::live_adapter::LiveAdapterStatus {
meerkat_core::live_adapter::LiveAdapterStatus::Ready
}
async fn close(&self) -> Result<(), meerkat_core::live_adapter::LiveAdapterError> {
Ok(())
}
}
#[tokio::test]
async fn websocket_closes_on_terminal_observation_with_typed_reason() {
use meerkat_core::live_adapter::{LiveAdapterErrorCode, LiveAdapterObservation};
let host = Arc::new(LiveAdapterHost::new(Arc::new(NoOpProjectionSink)));
let session_id = meerkat_core::types::SessionId::new();
let channel_id = host.open_channel(session_id).await.unwrap();
host.attach_adapter(
&channel_id,
Arc::new(ScriptedAdapter::new(LiveAdapterObservation::Error {
code: LiveAdapterErrorCode::ProviderError,
message: "scripted terminal failure".into(),
})),
)
.await
.unwrap();
let state = Arc::new(LiveWsState::new(host));
let token = state.mint_token(channel_id.clone()).await;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let ws_state = Arc::clone(&state);
let server_handle =
tokio::spawn(async move { serve_live_ws_listener(listener, ws_state).await });
let url = format!("ws://{addr}{LIVE_WS_PATH}?token={token}&channel={channel_id}");
let (ws_stream, _) = tokio_tungstenite::connect_async(&url).await.unwrap();
use futures::StreamExt;
use tokio_tungstenite::tungstenite::Message;
let (_write, mut read) = ws_stream.split();
let mut saw_observation = false;
let mut saw_close_with_terminal_reason = false;
while let Some(msg) = read.next().await {
match msg {
Ok(Message::Text(text)) => {
if text.contains("\"observation\"") && text.contains("provider_error") {
saw_observation = true;
}
}
Ok(Message::Close(Some(frame))) => {
assert!(
frame.reason.contains("terminal:provider_error"),
"unexpected close reason: {}",
frame.reason
);
saw_close_with_terminal_reason = true;
break;
}
Ok(_) => continue,
Err(_) => break,
}
}
assert!(
saw_observation,
"client should have received the terminal observation JSON before the close frame"
);
assert!(
saw_close_with_terminal_reason,
"WS pump must close with a typed terminal:<code> reason on Error observations"
);
server_handle.abort();
}
#[tokio::test]
async fn websocket_forwards_command_rejected_without_closing() {
use meerkat_core::live_adapter::{LiveAdapterErrorCode, LiveAdapterObservation};
let host = Arc::new(LiveAdapterHost::new(Arc::new(NoOpProjectionSink)));
let session_id = meerkat_core::types::SessionId::new();
let channel_id = host.open_channel(session_id).await.unwrap();
host.attach_adapter(
&channel_id,
Arc::new(ScriptedAdapter::new(
LiveAdapterObservation::CommandRejected {
code: LiveAdapterErrorCode::ConfigRejected {
reason: meerkat_core::live_adapter::LiveConfigRejectionReason::ImageInputNotImplemented,
},
message: "image_input_not_implemented".into(),
},
)),
)
.await
.unwrap();
let state = Arc::new(LiveWsState::new(host));
let token = state.mint_token(channel_id.clone()).await;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let ws_state = Arc::clone(&state);
let server_handle =
tokio::spawn(async move { serve_live_ws_listener(listener, ws_state).await });
let url = format!("ws://{addr}{LIVE_WS_PATH}?token={token}&channel={channel_id}");
let (ws_stream, _) = tokio_tungstenite::connect_async(&url).await.unwrap();
use futures::StreamExt;
use tokio_tungstenite::tungstenite::Message;
let (_write, mut read) = ws_stream.split();
let mut saw_command_rejected = false;
let mut saw_close = false;
let deadline = tokio::time::Instant::now() + std::time::Duration::from_millis(300);
while let Some(remaining) = deadline.checked_duration_since(tokio::time::Instant::now()) {
match tokio::time::timeout(remaining, read.next()).await {
Ok(Some(Ok(Message::Text(text)))) => {
if text.contains("command_rejected")
&& text.contains("image_input_not_implemented")
{
saw_command_rejected = true;
}
}
Ok(Some(Ok(Message::Close(_)))) => {
saw_close = true;
break;
}
Ok(Some(Ok(_))) => continue,
Ok(Some(Err(_))) | Ok(None) => break,
Err(_) => break, }
}
assert!(
saw_command_rejected,
"client must receive the typed CommandRejected observation JSON"
);
assert!(
!saw_close,
"R5-9: WS must NOT close after a CommandRejected observation"
);
server_handle.abort();
}
#[tokio::test]
async fn websocket_rejects_invalid_token() {
let host = Arc::new(LiveAdapterHost::new(Arc::new(NoOpProjectionSink)));
let state = Arc::new(LiveWsState::new(host));
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let ws_state = Arc::clone(&state);
let server_handle =
tokio::spawn(async move { serve_live_ws_listener(listener, ws_state).await });
let url = format!("ws://{addr}{LIVE_WS_PATH}?token=bogus&channel=does_not_exist");
let (ws_stream, _) = tokio_tungstenite::connect_async(&url).await.unwrap();
let (_write, mut read) = futures::StreamExt::split(ws_stream);
use futures::StreamExt;
if let Some(Ok(msg)) = read.next().await {
let _ = msg;
}
server_handle.abort();
}
#[tokio::test]
async fn websocket_rejects_token_with_wrong_channel() {
let host = Arc::new(LiveAdapterHost::new(Arc::new(NoOpProjectionSink)));
let session_a = meerkat_core::types::SessionId::new();
let session_b = meerkat_core::types::SessionId::new();
let channel_a = host.open_channel(session_a).await.unwrap();
let channel_b = host.open_channel(session_b).await.unwrap();
let state = Arc::new(LiveWsState::new(host));
let token = state.mint_token(channel_a.clone()).await;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let ws_state = Arc::clone(&state);
let server_handle =
tokio::spawn(async move { serve_live_ws_listener(listener, ws_state).await });
let url = format!("ws://{addr}{LIVE_WS_PATH}?token={token}&channel={channel_b}");
let (ws_stream, _) = tokio_tungstenite::connect_async(&url).await.unwrap();
use futures::StreamExt;
use tokio_tungstenite::tungstenite::Message;
let (_write, mut read) = ws_stream.split();
let mut saw_invalid_token = false;
while let Some(msg) = read.next().await {
match msg {
Ok(Message::Text(text)) => {
if text.contains("invalid_token") {
saw_invalid_token = true;
}
}
Ok(Message::Close(Some(frame))) => {
assert!(
frame.reason.contains("invalid_token"),
"unexpected close reason: {}",
frame.reason
);
saw_invalid_token = true;
break;
}
Ok(_) => continue,
Err(_) => break,
}
}
assert!(
saw_invalid_token,
"channel mismatch must surface as invalid_token to the client"
);
server_handle.abort();
}
#[tokio::test]
async fn websocket_rejects_binary_without_format() {
let host = Arc::new(LiveAdapterHost::new(Arc::new(NoOpProjectionSink)));
let session_id = meerkat_core::types::SessionId::new();
let channel_id = host.open_channel(session_id).await.unwrap();
host.attach_adapter(&channel_id, Arc::new(IdleAdapter))
.await
.unwrap();
let state = Arc::new(LiveWsState::new(host));
let token = state.mint_token(channel_id.clone()).await;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let ws_state = Arc::clone(&state);
let server_handle =
tokio::spawn(async move { serve_live_ws_listener(listener, ws_state).await });
let url = format!("ws://{addr}{LIVE_WS_PATH}?token={token}&channel={channel_id}");
let (ws_stream, _) = tokio_tungstenite::connect_async(&url).await.unwrap();
use futures::{SinkExt, StreamExt};
use tokio_tungstenite::tungstenite::Message;
let (mut write, mut read) = ws_stream.split();
write
.send(Message::Binary(vec![0u8; 32].into()))
.await
.unwrap();
let mut saw_close = false;
while let Some(msg) = read.next().await {
match msg {
Ok(Message::Close(Some(frame))) => {
assert!(
frame.reason.contains("binary_format_unnegotiated"),
"unexpected close reason: {}",
frame.reason
);
saw_close = true;
break;
}
Ok(_) => continue,
Err(_) => break,
}
}
assert!(saw_close, "expected close frame for un-negotiated binary");
server_handle.abort();
}
#[tokio::test]
async fn websocket_closes_on_invalid_text_frame() {
let host = Arc::new(LiveAdapterHost::new(Arc::new(NoOpProjectionSink)));
let session_id = meerkat_core::types::SessionId::new();
let channel_id = host.open_channel(session_id).await.unwrap();
host.attach_adapter(&channel_id, Arc::new(IdleAdapter))
.await
.unwrap();
let state = Arc::new(LiveWsState::new(host));
let token = state.mint_token(channel_id.clone()).await;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let ws_state = Arc::clone(&state);
let server_handle =
tokio::spawn(async move { serve_live_ws_listener(listener, ws_state).await });
let url = format!("ws://{addr}{LIVE_WS_PATH}?token={token}&channel={channel_id}");
let (ws_stream, _) = tokio_tungstenite::connect_async(&url).await.unwrap();
use futures::{SinkExt, StreamExt};
use tokio_tungstenite::tungstenite::Message;
let (mut write, mut read) = ws_stream.split();
write.send(Message::Text("not json".into())).await.unwrap();
let mut saw_close = false;
while let Some(msg) = read.next().await {
match msg {
Ok(Message::Close(Some(frame))) => {
assert!(
frame.reason.contains("invalid_frame"),
"unexpected close reason: {}",
frame.reason
);
saw_close = true;
break;
}
Ok(_) => continue,
Err(_) => break,
}
}
assert!(saw_close, "expected close frame for invalid JSON");
server_handle.abort();
}
struct DelayedObservationAdapter {
observation: tokio::sync::Mutex<Option<meerkat_core::live_adapter::LiveAdapterObservation>>,
delay: Duration,
}
impl DelayedObservationAdapter {
fn new(
observation: meerkat_core::live_adapter::LiveAdapterObservation,
delay: Duration,
) -> Self {
Self {
observation: tokio::sync::Mutex::new(Some(observation)),
delay,
}
}
}
#[async_trait::async_trait]
impl meerkat_core::live_adapter::LiveAdapter for DelayedObservationAdapter {
async fn send_command(
&self,
_command: meerkat_core::live_adapter::LiveAdapterCommand,
) -> Result<(), meerkat_core::live_adapter::LiveAdapterError> {
Ok(())
}
async fn next_observation(
&self,
) -> Result<
Option<meerkat_core::live_adapter::LiveAdapterObservation>,
meerkat_core::live_adapter::LiveAdapterError,
> {
tokio::time::sleep(self.delay).await;
let mut slot = self.observation.lock().await;
if let Some(obs) = slot.take() {
Ok(Some(obs))
} else {
std::future::pending().await
}
}
fn status(&self) -> meerkat_core::live_adapter::LiveAdapterStatus {
meerkat_core::live_adapter::LiveAdapterStatus::Ready
}
async fn close(&self) -> Result<(), meerkat_core::live_adapter::LiveAdapterError> {
Ok(())
}
}
#[tokio::test]
async fn observation_arm_not_starved_by_saturating_mic_audio() {
use meerkat_core::live_adapter::LiveAdapterObservation;
use meerkat_core::types::{StopReason, Usage};
let host = Arc::new(LiveAdapterHost::new(Arc::new(NoOpProjectionSink)));
let session_id = meerkat_core::types::SessionId::new();
let channel_id = host.open_channel(session_id).await.unwrap();
host.attach_adapter(
&channel_id,
Arc::new(DelayedObservationAdapter::new(
LiveAdapterObservation::TurnCompleted {
response_id: Some("resp_1".into()),
stop_reason: StopReason::EndTurn,
usage: Usage::default(),
},
Duration::from_millis(50),
)),
)
.await
.unwrap();
host.apply_status_update(
&channel_id,
meerkat_core::live_adapter::LiveAdapterStatus::Ready,
)
.await
.unwrap();
let state = Arc::new(LiveWsState::new(host));
let token = state.mint_token(channel_id.clone()).await;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let ws_state = Arc::clone(&state);
let server_handle =
tokio::spawn(async move { serve_live_ws_listener(listener, ws_state).await });
let url = format!(
"ws://{addr}{LIVE_WS_PATH}?token={token}&channel={channel_id}&format=pcm_24k_mono"
);
let (ws_stream, _) = tokio_tungstenite::connect_async(&url).await.unwrap();
use futures::{SinkExt, StreamExt};
use tokio_tungstenite::tungstenite::Message;
let (mut write, mut read) = ws_stream.split();
let saturator = tokio::spawn(async move {
let frame = vec![0u8; 960]; loop {
if write
.send(Message::Binary(frame.clone().into()))
.await
.is_err()
{
break;
}
tokio::task::yield_now().await;
}
});
let deadline = tokio::time::Instant::now() + Duration::from_millis(1500);
let mut saw_turn_completed = false;
while let Some(remaining) = deadline.checked_duration_since(tokio::time::Instant::now()) {
match tokio::time::timeout(remaining, read.next()).await {
Ok(Some(Ok(Message::Text(text)))) => {
if text.contains("turn_completed") || text.contains("\"resp_1\"") {
saw_turn_completed = true;
break;
}
}
Ok(Some(Ok(_))) => continue,
Ok(Some(Err(_))) | Ok(None) => break,
Err(_) => break,
}
}
saturator.abort();
server_handle.abort();
assert!(
saw_turn_completed,
"observation arm starved by mic-audio saturation — biased ordering or unpinned observation future regression"
);
}
#[tokio::test]
async fn websocket_forwards_observation_through_wire_mirror() {
use meerkat_contracts::WireLiveAdapterObservation;
use meerkat_core::live_adapter::LiveAdapterObservation;
let host = Arc::new(LiveAdapterHost::new(Arc::new(NoOpProjectionSink)));
let session_id = meerkat_core::types::SessionId::new();
let channel_id = host.open_channel(session_id).await.unwrap();
host.attach_adapter(
&channel_id,
Arc::new(ScriptedAdapter::new(LiveAdapterObservation::Ready)),
)
.await
.unwrap();
host.apply_status_update(
&channel_id,
meerkat_core::live_adapter::LiveAdapterStatus::Ready,
)
.await
.unwrap();
let state = Arc::new(LiveWsState::new(host));
let token = state.mint_token(channel_id.clone()).await;
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
let addr = listener.local_addr().unwrap();
let ws_state = Arc::clone(&state);
let server_handle =
tokio::spawn(async move { serve_live_ws_listener(listener, ws_state).await });
let url = format!("ws://{addr}{LIVE_WS_PATH}?token={token}&channel={channel_id}");
let (ws_stream, _) = tokio_tungstenite::connect_async(&url).await.unwrap();
use futures::StreamExt;
use tokio_tungstenite::tungstenite::Message;
let (_write, mut read) = ws_stream.split();
let deadline = tokio::time::Instant::now() + Duration::from_millis(1500);
let mut typed_observation = None;
while let Some(remaining) = deadline.checked_duration_since(tokio::time::Instant::now()) {
match tokio::time::timeout(remaining, read.next()).await {
Ok(Some(Ok(Message::Text(text)))) => {
match serde_json::from_str::<WireLiveAdapterObservation>(&text) {
Ok(obs) => {
typed_observation = Some(obs);
break;
}
Err(err) => panic!(
"WS forwarded JSON must deserialize as WireLiveAdapterObservation; got error {err} on payload: {text}"
),
}
}
Ok(Some(Ok(_))) => continue,
Ok(Some(Err(_))) | Ok(None) => break,
Err(_) => break,
}
}
server_handle.abort();
match typed_observation {
Some(WireLiveAdapterObservation::Ready) => {}
other => panic!("expected wire-mirror Ready, got {other:?}"),
}
}
}