use std::collections::BTreeSet;
use serde_json::{Value, json};
use crate::completion::{CacheRetention, Effort, Reasoning, ServiceTier, Verbosity};
use crate::completion::{
CompletionRequest, CompletionResponse, GenerationOptions, OnUnsupported, ProviderExtension,
ProviderOptions, ToolDefinition,
};
use crate::error::ProviderError;
use crate::message::{ToolChoice, ToolName};
use crate::operation::Completion;
use crate::test_utils::{RecordingHttpClient, json_body};
use crate::wire::{Encoded, Mode, Operation, Wire};
pub(crate) fn encoded_body<W>(
wire: &W,
request: CompletionRequest,
mode: Mode,
) -> Result<Value, ProviderError>
where
W: Wire<Op = Completion, Payload = Encoded>,
{
let request = Completion::prepare(request, &wire.describe())?;
let encoded = wire.encode(request, mode)?;
Ok(json_body(&encoded.request))
}
pub(crate) fn request_with<P: ProviderExtension>(options: &P::Options) -> CompletionRequest {
let options = ProviderOptions::new()
.with::<P>(options)
.unwrap_or_else(|error| panic!("the options are sections: {error}"));
CompletionRequest::new("hi").provider_options(options)
}
pub(crate) fn body_with<P, W>(wire: &W, options: &P::Options) -> Value
where
P: ProviderExtension,
W: Wire<Op = Completion, Payload = Encoded>,
{
encoded_body(wire, request_with::<P>(options), Mode::Unary)
.unwrap_or_else(|error| panic!("the request encodes: {error}"))
}
pub(crate) fn recorded_reply(provider: &str, scenario: &str, index: usize) -> String {
let file = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../rig-cassette/fixtures/cassettes")
.join(provider)
.join(format!("{scenario}.yaml"));
let text = std::fs::read_to_string(&file)
.unwrap_or_else(|error| panic!("{} should be readable: {error}", file.display()));
let interaction = text
.split("\n---\n")
.nth(index)
.unwrap_or_else(|| panic!("{} should record interaction {index}", file.display()));
let reply = interaction
.split_once("\nthen:")
.unwrap_or_else(|| panic!("{} should record a reply", file.display()))
.1;
let body = reply
.split_once(" body: '")
.unwrap_or_else(|| panic!("{} should record a single-quoted body", file.display()))
.1
.lines()
.next()
.unwrap_or_default()
.trim_end()
.trim_end_matches('\'');
body.replace("''", "'")
}
pub(crate) async fn reply_of<W>(wire: W, body: impl Into<String>) -> CompletionResponse
where
W: Wire<Op = Completion>,
RecordingHttpClient: crate::driver::Transport<W>,
{
let body: String = body.into();
crate::driver::Model::new(wire, RecordingHttpClient::new(bytes::Bytes::from(body)))
.call(CompletionRequest::new("hi"))
.await
.unwrap_or_else(|error| panic!("the reply decodes: {error}"))
}
pub(crate) fn recorded_stream(provider: &str, scenario: &str, index: usize) -> String {
let file = std::path::Path::new(env!("CARGO_MANIFEST_DIR"))
.join("../rig-cassette/fixtures/cassettes")
.join(provider)
.join(format!("{scenario}.yaml"));
let text = std::fs::read_to_string(&file)
.unwrap_or_else(|error| panic!("{} should be readable: {error}", file.display()));
let interaction = text
.split("\n---\n")
.nth(index)
.unwrap_or_else(|| panic!("{} should record interaction {index}", file.display()));
let reply = interaction
.split_once("\nthen:")
.unwrap_or_else(|| panic!("{} should record a reply", file.display()))
.1;
let block = reply
.split_once(" body: |+\n")
.unwrap_or_else(|| panic!("{} should record an event stream", file.display()))
.1;
let mut body = String::new();
for line in block.lines() {
if !line.is_empty() && !line.starts_with(" ") {
break;
}
body.push_str(line.get(4..).unwrap_or_default());
body.push('\n');
}
body
}
pub(crate) async fn streamed_reply_of<W>(wire: W, body: impl Into<String>) -> CompletionResponse
where
W: Wire<Op = Completion>,
crate::test_utils::MockStreamingClient: crate::driver::Transport<W>,
{
use futures::StreamExt;
let body: String = body.into();
let model = crate::driver::Model::new(
wire,
crate::test_utils::MockStreamingClient {
sse_bytes: bytes::Bytes::from(body),
},
);
let mut stream = model
.stream(CompletionRequest::new("hi"))
.unwrap_or_else(|error| panic!("the stream opens: {error}"));
while let Some(item) = stream.next().await {
item.unwrap_or_else(|error| panic!("the stream yields no error: {error}"));
}
stream
.finish()
.await
.unwrap_or_else(|error| panic!("the stream ends: {error}"))
}
pub(crate) fn chat_stream(chunks: &[Value]) -> String {
let mut body: String = chunks
.iter()
.map(|chunk| format!("data: {chunk}\n\n"))
.collect();
body.push_str("data: [DONE]\n\n");
body
}
pub(crate) fn chat_reply(extra: Value) -> String {
let mut reply = json!({
"id": "reply-1",
"object": "chat.completion",
"created": 0,
"model": "model",
"choices": [{
"index": 0,
"finish_reason": "stop",
"message": {"role": "assistant", "content": "pong"}
}],
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2}
});
merge(&mut reply, extra);
reply.to_string()
}
fn merge(into: &mut Value, over: Value) {
match (into, over) {
(Value::Object(into), Value::Object(over)) => {
for (key, value) in over {
match into.get_mut(&key) {
Some(slot) => merge(slot, value),
None => {
into.insert(key, value);
}
}
}
}
(Value::Array(into), Value::Array(over)) => {
for (index, value) in over.into_iter().enumerate() {
match into.get_mut(index) {
Some(slot) => merge(slot, value),
None => into.push(value),
}
}
}
(into, over) => *into = over,
}
}
fn option_variants() -> Vec<GenerationOptions> {
let mut variants: Vec<GenerationOptions> = [
Reasoning::Off,
Reasoning::Budget { tokens: 1024 },
Effort::Minimal.into(),
Effort::Low.into(),
Effort::Medium.into(),
Effort::High.into(),
Effort::XHigh.into(),
Effort::Max.into(),
]
.into_iter()
.map(|reasoning| GenerationOptions::default().reasoning(reasoning))
.collect();
for cache in [
CacheRetention::None,
CacheRetention::Short,
CacheRetention::Long,
] {
variants.push(GenerationOptions::default().cache(cache));
}
for tier in [
ServiceTier::Auto,
ServiceTier::Default,
ServiceTier::Flex,
ServiceTier::Priority,
] {
variants.push(GenerationOptions::default().service_tier(tier));
}
variants.push(GenerationOptions::default().verbosity(Verbosity::Low));
variants.push(GenerationOptions::default().parallel_tool_calls(true));
variants.push(GenerationOptions::default().top_p(0.5));
variants.push(GenerationOptions::default().seed(7));
variants.push(GenerationOptions::default().stop(["END"]));
variants
}
fn full_request(options: GenerationOptions) -> CompletionRequest {
let tool = ToolDefinition::new(
ToolName::new("lookup").unwrap_or_else(|error| panic!("{error}")),
"look a thing up",
json!({"type": "object", "properties": {"q": {"type": "string"}}}),
);
CompletionRequest::new("hi")
.preamble("be brief")
.tool(tool)
.tool_choice(ToolChoice::Required)
.temperature(0.5)
.max_tokens(8192)
.options(options.on_unsupported(OnUnsupported::Ignore))
}
fn leaves(value: &Value, at: &str, out: &mut BTreeSet<String>) {
match value {
Value::Object(fields) if !fields.is_empty() => {
for (key, field) in fields {
let key = key.replace('~', "~0").replace('/', "~1");
leaves(field, &format!("{at}/{key}"), out);
}
}
_ => {
out.insert(at.to_owned());
}
}
}
fn reserved_leaves<W>(wires: &[W]) -> BTreeSet<String>
where
W: Wire<Op = Completion, Payload = Encoded>,
{
let mut reserved = BTreeSet::new();
let mut requests = vec![full_request(GenerationOptions::default())];
requests.extend(option_variants().into_iter().map(full_request));
for wire in wires {
for request in &requests {
for mode in [Mode::Unary, Mode::Streaming] {
if let Ok(body) = encoded_body(wire, request.clone(), mode) {
leaves(&body, "", &mut reserved);
}
}
}
}
reserved
}
pub(crate) fn assert_no_reserved_leaf<P, W>(wires: &[W], options: &P::Options)
where
P: ProviderExtension,
W: Wire<Op = Completion, Payload = Encoded>,
{
let reserved = reserved_leaves(wires);
let sections = serde_json::to_value(options).unwrap_or_else(|error| panic!("{error}"));
let Value::Object(sections) = sections else {
panic!("{} options are an object of sections", P::PROVIDER);
};
assert!(!sections.is_empty(), "the options set no field");
for (section, fields) in §ions {
if fields.as_object().is_some_and(|fields| fields.is_empty()) {
continue;
}
let mut written = BTreeSet::new();
leaves(fields, "", &mut written);
for leaf in &written {
let clash = reserved.iter().find(|owned| {
*owned == leaf
|| owned.starts_with(&format!("{leaf}/"))
|| leaf.starts_with(&format!("{owned}/"))
});
assert!(
clash.is_none(),
"{}.{section} writes `{leaf}`, which `{}` owns",
P::PROVIDER,
clash.map(String::as_str).unwrap_or_default()
);
}
}
}