use futures::StreamExt;
use serde_json::{Value, json};
use super::DynModel;
use crate::completion::{CompletionRequest, Usage};
use crate::message::AssistantContent;
use crate::operation::Completion;
use crate::streaming::StreamEvent;
use crate::test_utils::{
CapturedSpan, MockCompletionModel, MockStreamEvent, MockTurn, TraceCapture,
};
fn spans_of(body: impl FnOnce()) -> Vec<Value> {
let capture = TraceCapture::default();
tracing::subscriber::with_default(capture.subscriber(), body);
capture.spans().iter().map(CapturedSpan::summary).collect()
}
fn request() -> CompletionRequest {
CompletionRequest::new("hello")
.preamble("be brief")
.model("probe-model")
}
fn usage() -> Usage {
Usage {
input_tokens: Some(7),
output_tokens: Some(3),
total_tokens: Some(10),
..Usage::default()
}
}
fn unary_turn() -> MockTurn {
MockTurn::tool_call("call_1", "lookup", json!({"q": 1}))
.with_response_id("resp_1")
.with_provider_request_id("req_1")
.with_usage(usage())
}
fn stream_turn() -> Vec<MockStreamEvent> {
vec![
MockStreamEvent::text("hel"),
MockStreamEvent::text("lo"),
MockStreamEvent::tool_call("call_1", "lookup", json!({"q": 1})),
MockStreamEvent::final_response(usage()),
]
}
#[test]
fn an_erased_unary_call_matches_the_direct_call_and_its_span() {
let _isolation = crate::test_utils::scoped_tracing_subscriber_guard_blocking();
let model = MockCompletionModel::from_turns([unary_turn(), unary_turn()]);
let dyn_model: DynModel<Completion> = model.clone().erase();
let mut direct = None;
let direct_spans = spans_of(|| {
direct = Some(futures::executor::block_on(model.call(request())));
});
let mut erased = None;
let erased_spans = spans_of(|| {
erased = Some(futures::executor::block_on(dyn_model.call(request())));
});
let direct = direct.expect("ran").expect("the direct call succeeds");
let erased = erased.expect("ran").expect("the erased call succeeds");
assert_eq!(
serde_json::to_value(&direct).expect("json"),
serde_json::to_value(&erased).expect("json"),
"the erased call folds the same response"
);
assert_eq!(direct.provider_request_id.as_deref(), Some("req_1"));
assert!(!direct_spans.is_empty(), "the direct call opened a span");
assert_eq!(
direct_spans, erased_spans,
"the erased call records the same span"
);
}
#[test]
fn an_erased_stream_yields_the_direct_stream_item_for_item() {
let _isolation = crate::test_utils::scoped_tracing_subscriber_guard_blocking();
let model = MockCompletionModel::from_stream_turns([stream_turn(), stream_turn()]);
let dyn_model = DynModel::from(model.clone());
let mut direct = Vec::new();
let direct_spans = spans_of(|| {
let stream = model.stream(request()).expect("the direct stream opens");
direct = futures::executor::block_on(stream.collect::<Vec<_>>());
});
let mut erased = Vec::new();
let erased_spans = spans_of(|| {
let stream = dyn_model
.stream(request())
.expect("the erased stream opens");
erased = futures::executor::block_on(stream.collect::<Vec<_>>());
});
assert!(
direct.iter().any(|item| matches!(
item,
Ok(crate::streaming::Item::Event(StreamEvent::End {
content: AssistantContent::ToolCall(_),
..
}))
)),
"a tool call end carries its finalized call: {direct:?}"
);
assert_eq!(
format!("{direct:?}"),
format!("{erased:?}"),
"the erased stream yields the same items"
);
assert_eq!(
direct_spans, erased_spans,
"the erased stream records the same span"
);
}
#[test]
fn an_erased_model_names_its_wire() {
let erased = MockCompletionModel::text("x").erase();
assert_eq!(erased.name(), crate::test_utils::MOCK_PROVIDER);
assert_eq!(erased.id(), None);
assert_eq!(
format!("{erased:?}"),
r#"DynModel { name: "mock", id: None }"#
);
let clone = erased.clone();
assert_eq!(clone.name(), erased.name());
}
#[test]
fn an_erased_call_on_a_borrowed_prompt_outlives_the_prompt() {
fn spawnable<F: std::future::Future + Send + 'static>(future: F) -> F {
future
}
let dyn_model: DynModel<Completion> =
MockCompletionModel::from_turns([MockTurn::text("hi")]).erase();
let future = {
let prompt = String::from("Say hi.");
spawnable(dyn_model.call(prompt.as_str()))
};
let response = futures::executor::block_on(future).expect("the call succeeds");
assert_eq!(response.text(), "hi");
}