#![allow(dead_code)]
use rig_core::driver::Model;
use rig_core::http_client;
use rig_core::providers::openai::OpenAIConfig;
use rig_core::providers::openai::responses_api::websocket::ResponsesWebSocketSession;
use rig_core::providers::openai::responses_api::wire::Responses;
use rig_core::test_utils::RecordingHttpClient;
use rig_core::wasm_compat::WasmBoxedFuture;
use rig_core::ws_client::{BoxedWebSocketConnection, CloseFrame, Frame, WebSocketConnection};
use std::collections::VecDeque;
use std::sync::{Arc, Mutex};
use std::time::Duration;
pub type TestClient = Model<Responses>;
pub type TestSession = ResponsesWebSocketSession;
#[derive(Clone, Copy, Default, PartialEq, Eq, Debug)]
pub enum WhenDrained {
#[default]
EndStream,
Stall,
}
#[derive(Default)]
struct ScriptState {
turns: VecDeque<Vec<Frame>>,
inbound: VecDeque<Frame>,
sent: Vec<String>,
closed: bool,
drained: WhenDrained,
}
#[derive(Clone, Default)]
pub struct Script(Arc<Mutex<ScriptState>>);
impl Script {
pub fn turns<I, J>(turns: I) -> Self
where
I: IntoIterator<Item = J>,
J: IntoIterator<Item = String>,
{
let state = ScriptState {
turns: turns
.into_iter()
.map(|turn| turn.into_iter().map(Frame::Text).collect())
.collect(),
..ScriptState::default()
};
Self(Arc::new(Mutex::new(state)))
}
pub fn turn<I: IntoIterator<Item = String>>(frames: I) -> Self {
Self::turns([frames])
}
#[must_use]
pub fn stalling(self) -> Self {
self.0.lock().expect("script lock").drained = WhenDrained::Stall;
self
}
pub fn sent(&self) -> Vec<String> {
self.0.lock().expect("script lock").sent.clone()
}
pub fn closed(&self) -> bool {
self.0.lock().expect("script lock").closed
}
pub fn connection(&self) -> BoxedWebSocketConnection {
Box::new(ScriptedConnection(self.clone()))
}
}
struct ScriptedConnection(Script);
impl WebSocketConnection for ScriptedConnection {
fn send(&mut self, frame: Frame) -> WasmBoxedFuture<'_, http_client::Result<()>> {
let mut state = self.0.0.lock().expect("script lock");
match frame {
Frame::Text(text) => state.sent.push(text),
other => panic!("the session only writes text frames, got {other:?}"),
}
if let Some(turn) = state.turns.pop_front() {
state.inbound.extend(turn);
}
Box::pin(std::future::ready(Ok(())))
}
fn recv(&mut self) -> WasmBoxedFuture<'_, http_client::Result<Option<Frame>>> {
let next = {
let mut state = self.0.0.lock().expect("script lock");
match state.inbound.pop_front() {
Some(frame) => Some(Some(frame)),
None => match state.drained {
WhenDrained::EndStream => Some(None),
WhenDrained::Stall => None,
},
}
};
match next {
Some(frame) => Box::pin(std::future::ready(Ok(frame))),
None => Box::pin(std::future::pending()),
}
}
fn close(
&mut self,
_frame: Option<CloseFrame>,
) -> WasmBoxedFuture<'_, http_client::Result<()>> {
self.0.0.lock().expect("script lock").closed = true;
Box::pin(std::future::ready(Ok(())))
}
}
pub fn test_client() -> TestClient {
OpenAIConfig::new("test-key")
.connect(RecordingHttpClient::new("{}"))
.responses("gpt-4o")
}
pub fn session(client: &TestClient, script: &Script) -> TestSession {
session_with_timeout(client, script, None)
}
pub fn session_with_timeout(
client: &TestClient,
script: &Script,
event_timeout: Option<Duration>,
) -> TestSession {
ResponsesWebSocketSession::from_connection(
client.wire.clone(),
script.connection(),
event_timeout,
)
}