use serde::de::DeserializeOwned;
use thiserror::Error;
use rig_core::{
memory::MemoryError,
wasm_compat::{WasmCompatSend, WasmCompatSync},
};
pub use rig_core::completion::*;
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum PromptError {
#[error("CompletionError: {0}")]
CompletionError(#[from] CompletionError),
#[error("MemoryError: {0}")]
MemoryError(#[from] MemoryError),
#[error("MaxTurnsError: reached max turns limit: {max_turns}")]
MaxTurnsError {
max_turns: usize,
chat_history: Box<Vec<Message>>,
prompt: Box<Message>,
},
#[error("PromptCancelled: {reason}")]
PromptCancelled {
chat_history: Vec<Message>,
reason: String,
},
#[error(
"UnknownToolCall: model attempted to call unknown or disallowed tool `{tool_name}`. Available tools: {available_tools:?}. Allowed tools for this turn: {allowed_tools:?}"
)]
UnknownToolCall {
tool_name: String,
available_tools: Vec<String>,
allowed_tools: Vec<String>,
chat_history: Box<Vec<Message>>,
},
}
impl PromptError {
pub fn provider_response_body(&self) -> Option<&str> {
match self {
Self::CompletionError(error) => error.provider_response_body(),
_ => None,
}
}
pub fn provider_response_json(&self) -> Result<Option<serde_json::Value>, serde_json::Error> {
match self {
Self::CompletionError(error) => error.provider_response_json(),
_ => Ok(None),
}
}
pub fn provider_response_status(&self) -> Option<http::StatusCode> {
match self {
Self::CompletionError(error) => error.provider_response_status(),
_ => None,
}
}
pub(crate) fn prompt_cancelled(
chat_history: impl IntoIterator<Item = Message>,
reason: impl Into<String>,
) -> Self {
Self::PromptCancelled {
chat_history: chat_history.into_iter().collect(),
reason: reason.into(),
}
}
}
#[derive(Debug, Error)]
#[non_exhaustive]
pub enum StructuredOutputError {
#[error("PromptError: {0}")]
PromptError(#[from] Box<PromptError>),
#[error("DeserializationError: {0}")]
DeserializationError(#[from] serde_json::Error),
#[error("EmptyResponse: model returned no content")]
EmptyResponse,
}
impl StructuredOutputError {
pub fn provider_response_body(&self) -> Option<&str> {
match self {
Self::PromptError(error) => error.provider_response_body(),
_ => None,
}
}
pub fn provider_response_json(&self) -> Result<Option<serde_json::Value>, serde_json::Error> {
match self {
Self::PromptError(error) => error.provider_response_json(),
_ => Ok(None),
}
}
pub fn provider_response_status(&self) -> Option<http::StatusCode> {
match self {
Self::PromptError(error) => error.provider_response_status(),
_ => None,
}
}
}
pub trait Prompt: WasmCompatSend + WasmCompatSync {
fn prompt(
&self,
prompt: impl Into<Message> + WasmCompatSend,
) -> impl std::future::IntoFuture<Output = Result<String, PromptError>, IntoFuture: WasmCompatSend>;
}
pub trait Chat: WasmCompatSend + WasmCompatSync {
fn chat(
&self,
prompt: impl Into<Message> + WasmCompatSend,
chat_history: &mut Vec<Message>,
) -> impl std::future::Future<Output = Result<String, PromptError>> + WasmCompatSend;
}
pub trait TypedPrompt: WasmCompatSend + WasmCompatSync {
type TypedRequest<T>: std::future::IntoFuture<Output = Result<T, StructuredOutputError>>
where
T: schemars::JsonSchema + DeserializeOwned + WasmCompatSend + 'static;
fn prompt_typed<T>(&self, prompt: impl Into<Message> + WasmCompatSend) -> Self::TypedRequest<T>
where
T: schemars::JsonSchema + DeserializeOwned + WasmCompatSend;
}
#[cfg(test)]
mod provider_response_tests {
use rig_core::{ProviderResponseError, http_client};
use super::*;
#[test]
fn prompt_error_forwards_provider_response_to_completion_error() {
let body = r#"{"error":{"message":"boom"}}"#;
let inner =
CompletionError::from_http_response(http::StatusCode::SERVICE_UNAVAILABLE, body);
let error = PromptError::CompletionError(inner);
assert_eq!(
error.provider_response_status(),
Some(http::StatusCode::SERVICE_UNAVAILABLE),
);
assert_eq!(error.provider_response_body(), Some(body));
assert_eq!(
error
.provider_response_json()
.expect("valid json")
.expect("present json")["error"]["message"],
"boom",
);
}
#[test]
fn prompt_error_provider_response_helpers_forward_http_status_and_body() {
let body = r#"{"error":{"message":"unauthorized"}}"#;
let error = PromptError::CompletionError(CompletionError::HttpError(
http_client::Error::InvalidStatusCodeWithMessage(
http::StatusCode::UNAUTHORIZED,
body.to_string(),
),
));
assert_eq!(error.provider_response_body(), Some(body));
assert_eq!(
error.provider_response_status(),
Some(http::StatusCode::UNAUTHORIZED)
);
assert_eq!(
error.provider_response_json().expect("valid JSON body"),
Some(serde_json::json!({
"error": { "message": "unauthorized" }
}))
);
}
#[test]
fn prompt_error_provider_response_helpers_forward_wrapped_completion_error() {
let body = r#"{"error":{"code":"invalid_request","message":"bad input"}}"#;
let error = PromptError::CompletionError(CompletionError::ProviderResponse(
ProviderResponseError {
status: None,
body: body.to_string(),
},
));
assert_eq!(error.provider_response_body(), Some(body));
assert_eq!(error.provider_response_status(), None);
assert_eq!(
error.provider_response_json().expect("valid JSON body"),
Some(serde_json::json!({
"error": {
"code": "invalid_request",
"message": "bad input"
}
}))
);
}
#[test]
fn prompt_error_provider_response_helpers_return_none_for_unrelated_variant() {
let error = PromptError::PromptCancelled {
chat_history: vec![Message::user("hi")],
reason: "cancelled".to_string(),
};
assert_eq!(error.provider_response_body(), None);
assert_eq!(error.provider_response_status(), None);
assert_eq!(
error
.provider_response_json()
.expect("no body is not an error"),
None
);
}
#[test]
fn structured_output_error_provider_response_helpers_forward_prompt_error() {
let body = r#"{"error":{"message":"bad input"}}"#;
let error = StructuredOutputError::PromptError(Box::new(PromptError::CompletionError(
CompletionError::ProviderResponse(ProviderResponseError {
status: Some(http::StatusCode::BAD_REQUEST),
body: body.to_string(),
}),
)));
assert_eq!(error.provider_response_body(), Some(body));
assert_eq!(
error.provider_response_status(),
Some(http::StatusCode::BAD_REQUEST)
);
}
}