use std::collections::BTreeMap;
use std::sync::Arc;
use std::time::Duration;
use ferrin_message::Message;
use ferrin_spec::FinishReason;
use ferrin_spec::JsonValue;
use ferrin_spec::ModelId;
use ferrin_spec::ProviderId;
use ferrin_spec::ResponseMetadata;
use ferrin_spec::ToolCallId;
use ferrin_spec::ToolName;
use ferrin_spec::Usage;
use ferrin_spec::Warning;
use ferrin_spec::language_model::CallOptionsRecord;
use serde::Deserialize;
use serde::Serialize;
use crate::error::Error;
use crate::generate_text::StepContent;
use crate::generate_text::StepPerformance;
use crate::generate_text::StepResult;
use crate::generate_text::ToolErrorInfo;
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ModelIdentity {
pub provider: ProviderId,
pub model_id: ModelId,
}
impl ModelIdentity {
#[must_use]
pub fn new(provider: impl Into<ProviderId>, model_id: impl Into<ModelId>) -> Self {
Self {
provider: provider.into(),
model_id: model_id.into(),
}
}
}
#[derive(Debug, Clone, PartialEq)]
pub struct RecordedInputs {
pub system: Option<crate::prompt::Instructions>,
pub messages: Arc<[Message]>,
}
#[derive(Debug, Clone)]
pub struct StartEvent {
pub runtime_context: Option<JsonValue>,
pub call_id: String,
pub function_id: Option<String>,
pub model: ModelIdentity,
pub inputs: Option<RecordedInputs>,
pub metadata: BTreeMap<String, JsonValue>,
}
#[derive(Debug, Clone)]
pub struct StepStartEvent {
pub runtime_context: Option<JsonValue>,
pub call_id: String,
pub step_number: u32,
pub model: ModelIdentity,
pub messages: Option<Arc<[Message]>>,
}
#[derive(Debug, Clone)]
pub struct ModelCallStartEvent {
pub runtime_context: Option<JsonValue>,
pub call_id: String,
pub step_number: u32,
pub model: ModelIdentity,
pub call_options: Option<CallOptionsRecord>,
}
#[derive(Debug, Clone)]
pub struct ModelCallEndEvent {
pub runtime_context: Option<JsonValue>,
pub call_id: String,
pub step_number: u32,
pub model: ModelIdentity,
pub content: Option<Vec<StepContent>>,
pub finish_reason: FinishReason,
pub usage: Usage,
pub response: ResponseMetadata,
pub performance: StepPerformance,
pub warnings: Vec<Warning>,
}
#[derive(Debug, Clone)]
pub struct ToolExecutionStartEvent {
pub runtime_context: Option<JsonValue>,
pub call_id: String,
pub tool_call_id: ToolCallId,
pub tool_name: ToolName,
pub input: Option<JsonValue>,
}
#[derive(Debug, Clone)]
pub struct ToolExecutionEndEvent {
pub runtime_context: Option<JsonValue>,
pub call_id: String,
pub tool_call_id: ToolCallId,
pub tool_name: ToolName,
pub output: Option<ToolOutcome>,
pub error: Option<ToolErrorInfo>,
pub duration: Duration,
}
#[derive(Debug, Clone, PartialEq)]
pub struct ToolOutcome {
pub output: JsonValue,
}
#[derive(Debug, Clone)]
pub struct StepEndEvent {
pub call_id: String,
pub step: Arc<StepResult>,
}
#[derive(Debug, Clone)]
pub struct EndEvent {
pub runtime_context: Option<JsonValue>,
pub call_id: String,
pub steps: Arc<[StepResult]>,
pub total_usage: Usage,
pub output_recorded: Option<JsonValue>,
}
#[derive(Debug, Clone)]
pub struct AbortEvent {
pub call_id: String,
pub steps_completed: u32,
}
#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)]
#[serde(rename_all = "kebab-case")]
#[non_exhaustive]
pub enum ErrorPhase {
Prompt,
ModelCall,
ToolExecution,
Output,
Stream,
}
#[derive(Debug)]
pub struct ErrorEvent<'a> {
pub call_id: &'a str,
pub error: &'a Error,
pub phase: ErrorPhase,
}
#[derive(Debug, Clone)]
pub struct EmbedStartEvent {
pub call_id: String,
pub model: ModelIdentity,
pub value_count: usize,
pub values: Option<Vec<String>>,
}
#[derive(Debug, Clone)]
pub struct EmbedEndEvent {
pub call_id: String,
pub embedding_count: usize,
pub tokens: Option<u64>,
pub duration: Duration,
}
#[derive(Debug, Clone)]
pub struct RerankStartEvent {
pub call_id: String,
pub model: ModelIdentity,
pub document_count: usize,
pub query: Option<String>,
}
#[derive(Debug, Clone)]
pub struct RerankEndEvent {
pub call_id: String,
pub ranked_count: usize,
pub duration: Duration,
}