mod seq;
pub mod tools;
pub mod transport;
use crate::{
protocol::{MSEvent, MSRequest, SeqGenReq, SeqOpenReq},
tools::Toolbox,
};
use futures::{Sink, Stream};
use futures_util::{SinkExt, StreamExt};
pub use seq::{EmbedOpts, EmbeddingResult, GenChunk, GenStream, Seq};
use std::{
collections::{HashMap, HashSet},
pin::Pin,
sync::Arc,
};
use thiserror::Error;
use tokio::sync::{mpsc, Mutex, RwLock};
use tokio_tungstenite::tungstenite::{protocol::Message, Error as WsError};
use tracing::{debug, error};
use uuid::Uuid;
#[derive(Error, Debug)]
pub enum ModelSocketError {
#[error("ws error: {0}")]
WebSocket(#[from] WsError),
#[error("invalid url: {0}")]
Url(#[from] url::ParseError),
#[error("json error: {0}")]
Json(#[from] serde_json::Error),
#[error("protocol error: {0}")]
Protocol(String),
#[error("state error: {0}")]
State(String),
#[error("open error: {0}")]
Open(String),
#[error("channel closed")]
Send(#[from] mpsc::error::SendError<Message>),
#[error("seq closed")]
SeqClosed,
#[error("command error: {0}")]
Command(String),
#[error("{message}")]
Remote {
message: String,
code: Option<String>,
details: Option<serde_json::Map<String, serde_json::Value>>,
},
#[error("{0}")]
Other(#[from] anyhow::Error),
}
#[derive(Clone)]
struct ModelSocketEventHandler {
opening_seqs: Arc<Mutex<HashMap<String, mpsc::Sender<Result<String, ModelSocketError>>>>>,
seqs: Arc<Mutex<HashMap<String, Seq>>>,
closed_seqs: Arc<RwLock<HashSet<String>>>,
}
#[derive(Clone)]
pub struct ModelSocket {
ws_sink: Arc<Mutex<Pin<Box<dyn Sink<MSRequest, Error = ModelSocketError> + Send>>>>,
opening_seqs: Arc<Mutex<HashMap<String, mpsc::Sender<Result<String, ModelSocketError>>>>>,
seqs: Arc<Mutex<HashMap<String, Seq>>>,
closed_seqs: Arc<RwLock<HashSet<String>>>,
}
impl ModelSocket {
pub fn new<RX, TX, T>(transport: T) -> Result<Self, ModelSocketError>
where
RX: Stream<Item = Result<MSEvent, ModelSocketError>> + Send + Unpin + 'static,
TX: Sink<MSRequest, Error = ModelSocketError> + Send + 'static,
T: transport::MSTransport<RX, TX>,
{
super::ensure_rustls_provider();
let (ws_sink, ws_stream) = transport.split();
let socket = Self {
ws_sink: Arc::new(Mutex::new(Box::pin(ws_sink))),
opening_seqs: Arc::new(Mutex::new(HashMap::new())),
seqs: Arc::new(Mutex::new(HashMap::new())),
closed_seqs: Default::default(),
};
let event_handler = socket.event_handler();
tokio::spawn(async move {
event_handler.event_recv_loop(ws_stream).await;
});
Ok(socket)
}
pub async fn connect(url: &str, api_key: Option<&str>) -> Result<Self, ModelSocketError> {
let ws_transport = transport::ws::WebSocketTransport::connect(url, api_key).await?;
Self::new(ws_transport)
}
fn clone_components(&self) -> Self {
Self {
ws_sink: self.ws_sink.clone(),
opening_seqs: self.opening_seqs.clone(),
seqs: self.seqs.clone(),
closed_seqs: self.closed_seqs.clone(),
}
}
fn event_handler(&self) -> ModelSocketEventHandler {
ModelSocketEventHandler {
opening_seqs: self.opening_seqs.clone(),
seqs: self.seqs.clone(),
closed_seqs: self.closed_seqs.clone(),
}
}
}
impl ModelSocketEventHandler {
async fn event_recv_loop<S: Stream<Item = Result<MSEvent, ModelSocketError>> + Unpin>(
self,
mut events: S,
) {
while let Some(event) = events.next().await {
match event {
Ok(event) => self.on_event(event).await,
Err(err) => {
error!("ms error: {}", err);
break;
}
}
}
}
async fn on_event(&self, event: MSEvent) {
debug!("<- {:?}", event);
match &event {
MSEvent::SeqOpened { seq_id, cid } => self.on_seq_opened(seq_id, cid).await,
MSEvent::Error {
cid,
seq_id,
message,
code,
details,
} => {
self.on_error(cid, message, code, details).await;
if seq_id.is_some() {
self.forward_to_seq(&event).await;
}
}
MSEvent::SeqClosed { seq_id, .. } => {
self.forward_to_seq(&event).await;
self.on_seq_closed(seq_id).await
}
_ => self.forward_to_seq(&event).await,
}
}
async fn is_seq_closed(&self, seq_id: &str) -> bool {
self.closed_seqs.read().await.contains(seq_id)
}
async fn forward_to_seq(&self, event: &MSEvent) {
let seq_id = match event.seq_id() {
Some(id) => id,
None => {
error!("unroutable event forwarded to seq: {:?}", event);
return;
}
};
let mut seqs = self.seqs.lock().await;
if self.is_seq_closed(seq_id).await {
debug!(seq = seq_id, "skipping event for closed seq");
return;
}
if let Some(seq) = seqs.get_mut(seq_id) {
seq.on_event(event).await;
} else {
error!("state error: unknown seq_id {}", seq_id);
}
}
async fn on_error(
&self,
cid: &Option<String>,
message: &str,
code: &Option<String>,
details: &Option<serde_json::Map<String, serde_json::Value>>,
) {
error!("error: {}", message);
if let Some(cid) = cid {
if let Some(sender) = self.opening_seqs.lock().await.remove(cid) {
let _ = sender.send(Err(remote_error(message, code, details))).await;
}
}
}
async fn on_seq_closed(&self, seq_id: &str) {
if let Some(_seq) = self.seqs.lock().await.remove(seq_id) {
debug!(seq = seq_id, "removing seq from client socket state");
self.closed_seqs.write().await.insert(seq_id.to_owned());
} else {
if !self.is_seq_closed(seq_id).await {
error!(seq_id, "unknown seq id during on close event");
}
}
}
async fn on_seq_opened(&self, seq_id: &str, cid: &str) {
let mut opening_seqs = self.opening_seqs.lock().await;
if let Some(sender) = opening_seqs.remove(cid) {
if let Err(e) = sender.send(Ok(seq_id.to_owned())).await {
error!("failed to send seq_id: {}", e);
}
} else {
error!("unknown opened seq cid {}", cid);
}
}
}
fn remote_error(
message: &str,
code: &Option<String>,
details: &Option<serde_json::Map<String, serde_json::Value>>,
) -> ModelSocketError {
ModelSocketError::Remote {
message: message.to_string(),
code: code.clone(),
details: details.clone(),
}
}
impl ModelSocket {
pub async fn open(&self, model: &str, opts: Option<OpenOpts>) -> Result<Seq, ModelSocketError> {
let cid = Uuid::new_v4().to_string();
let (tx, mut rx) = mpsc::channel(1);
{
let mut opening_seqs = self.opening_seqs.lock().await;
opening_seqs.insert(cid.clone(), tx);
}
let opts = opts.unwrap_or_default();
let tools_enabled = opts.toolbox.is_some();
self.send_request(MSRequest::SeqOpen {
cid: cid.clone(),
data: SeqOpenReq {
model: model.to_string(),
tools_enabled,
tool_prompt: opts.tool_prompt,
tool_schemas: opts.tool_schemas,
skip_prelude: opts.skip_prelude,
},
})
.await?;
let seq_id = rx.recv().await.unwrap_or_else(|| {
Err(ModelSocketError::Open(
"Failed to receive seq_id".to_string(),
))
})?;
let tool_def_prompt = opts.toolbox.as_ref().and_then(|t| t.tool_def_prompt());
let toolbox = Arc::new(opts.toolbox.map(|t| Mutex::new(t)));
let seq = Seq::new(
seq_id.clone(),
model.to_string(),
self.clone_components(),
toolbox,
None,
);
{
let mut seqs = self.seqs.lock().await;
let mut seq_event_handler = seq.clone();
seq_event_handler.event_tx = opts.event_sink;
seqs.insert(seq_id, seq_event_handler);
}
if let Some(tool_def_prompt) = tool_def_prompt {
seq.append(tool_def_prompt, AppendOpts::system()).await?;
}
Ok(seq)
}
async fn send_request(&self, req: MSRequest) -> Result<(), ModelSocketError> {
debug!("-> {:?}", req);
let mut sink = self.ws_sink.lock().await;
sink.send(req).await?;
Ok(())
}
}
#[derive(Default, Debug, Clone)]
pub struct AppendOpts {
pub role: Option<String>,
}
#[derive(Default, Debug, Clone)]
pub struct GenOpts {
pub role: Option<String>,
pub stop_strings: Option<Vec<String>>,
pub max_length: Option<u32>,
pub max_tokens: Option<u32>,
pub hidden: Option<bool>,
pub temperature: Option<f32>,
pub top_p: Option<f32>,
pub top_k: Option<i32>,
pub repeat_penalty: Option<f32>,
pub seed: Option<u64>,
pub json_schema: Option<serde_json::Value>,
pub json_schema_strict: Option<bool>,
pub frequency_penalty: Option<f32>,
pub presence_penalty: Option<f32>,
}
impl GenOpts {
pub fn assistant() -> Self {
Self {
role: Some("assistant".to_string()),
..Default::default()
}
}
pub fn user() -> Self {
Self {
role: Some("user".to_string()),
..Default::default()
}
}
pub fn system() -> Self {
Self {
role: Some("system".to_string()),
..Default::default()
}
}
}
impl Into<SeqGenReq> for GenOpts {
fn into(self) -> SeqGenReq {
SeqGenReq {
role: self.role,
stop_strings: self.stop_strings,
max_length: self.max_length,
max_tokens: self.max_tokens,
hidden: self.hidden.unwrap_or(false),
temperature: self.temperature,
top_p: self.top_p,
top_k: self.top_k,
repeat_penalty: self.repeat_penalty,
seed: self.seed,
json_schema: self.json_schema,
json_schema_strict: self.json_schema_strict,
frequency_penalty: self.frequency_penalty,
presence_penalty: self.presence_penalty,
..Default::default()
}
}
}
impl AppendOpts {
pub fn assistant() -> Self {
Self {
role: Some("assistant".to_string()),
}
}
pub fn user() -> Self {
Self {
role: Some("user".to_string()),
}
}
pub fn system() -> Self {
Self {
role: Some("system".to_string()),
}
}
}
#[derive(Default, Debug)]
pub struct OpenOpts {
pub tool_prompt: Option<String>,
pub skip_prelude: bool,
pub tool_schemas: Option<HashMap<String, serde_json::Value>>,
pub event_sink: Option<mpsc::Sender<Result<MSEvent, ModelSocketError>>>,
pub toolbox: Option<Box<dyn Toolbox>>,
}
#[cfg(test)]
mod tests {
use super::*;
use std::task::{Context, Poll};
use tokio::time::{timeout, Duration};
struct ChannelTransport {
requests: tokio::sync::mpsc::UnboundedSender<MSRequest>,
events: tokio::sync::mpsc::UnboundedReceiver<Result<MSEvent, ModelSocketError>>,
}
struct RequestSink(tokio::sync::mpsc::UnboundedSender<MSRequest>);
impl Sink<MSRequest> for RequestSink {
type Error = ModelSocketError;
fn poll_ready(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn start_send(self: Pin<&mut Self>, item: MSRequest) -> Result<(), Self::Error> {
self.0
.send(item)
.map_err(|_| ModelSocketError::State("request receiver closed".to_string()))
}
fn poll_flush(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
fn poll_close(
self: Pin<&mut Self>,
_cx: &mut Context<'_>,
) -> Poll<Result<(), Self::Error>> {
Poll::Ready(Ok(()))
}
}
struct EventStream(tokio::sync::mpsc::UnboundedReceiver<Result<MSEvent, ModelSocketError>>);
impl Stream for EventStream {
type Item = Result<MSEvent, ModelSocketError>;
fn poll_next(mut self: Pin<&mut Self>, cx: &mut Context<'_>) -> Poll<Option<Self::Item>> {
self.0.poll_recv(cx)
}
}
impl transport::MSTransport<EventStream, RequestSink> for ChannelTransport {
fn split(self) -> (RequestSink, EventStream) {
(RequestSink(self.requests), EventStream(self.events))
}
}
#[tokio::test]
async fn dropping_socket_closes_request_channel_while_event_stream_is_open() {
let (request_tx, mut request_rx) = tokio::sync::mpsc::unbounded_channel();
let (event_tx, event_rx) = tokio::sync::mpsc::unbounded_channel();
let socket = ModelSocket::new(ChannelTransport {
requests: request_tx,
events: event_rx,
})
.unwrap();
drop(socket);
let closed = timeout(Duration::from_millis(100), request_rx.recv())
.await
.expect("request receiver should observe closure after dropping ModelSocket");
assert!(closed.is_none());
drop(event_tx);
}
#[tokio::test]
async fn open_error_preserves_remote_details() {
let (request_tx, mut request_rx) = tokio::sync::mpsc::unbounded_channel();
let (event_tx, event_rx) = tokio::sync::mpsc::unbounded_channel();
let socket = ModelSocket::new(ChannelTransport {
requests: request_tx,
events: event_rx,
})
.unwrap();
let open = tokio::spawn(async move { socket.open("model", None).await });
let MSRequest::SeqOpen { cid, .. } = request_rx.recv().await.unwrap() else {
panic!("expected sequence open");
};
let details = serde_json::json!({
"rpm_limit": 60,
"rpm_remaining": 0,
"tpm_limit": 100_000,
"tpm_remaining": 0,
"retry_after_ms": 1_000
})
.as_object()
.unwrap()
.clone();
event_tx
.send(Ok(MSEvent::Error {
cid: Some(cid),
seq_id: Some("seq-1".into()),
message: "Rate limit exceeded".into(),
code: Some("rate_limit_exceeded".into()),
details: Some(details.clone()),
}))
.unwrap();
match open.await.unwrap() {
Err(ModelSocketError::Remote {
code,
details: actual,
..
}) => {
assert_eq!(code.as_deref(), Some("rate_limit_exceeded"));
assert_eq!(actual, Some(details));
}
_ => panic!("expected structured remote error"),
}
}
}