use chrono::{DateTime, Utc};
use serde::{Deserialize, Serialize};
use serde_json::Value;
use starweaver_core::{Metadata, RunId, SessionId, TraceContext};
use crate::records::ExecutionStatus;
pub struct ToolReturnRecordInput<'a> {
pub session_id: &'a SessionId,
pub run_id: &'a RunId,
pub tool_call_id: &'a str,
pub tool_name: &'a str,
pub metadata: &'a Metadata,
pub trace_context: Option<&'a TraceContext>,
pub policy: Option<Value>,
}
impl<'a> ToolReturnRecordInput<'a> {
#[must_use]
pub const fn new(
session_id: &'a SessionId,
run_id: &'a RunId,
tool_call_id: &'a str,
tool_name: &'a str,
metadata: &'a Metadata,
) -> Self {
Self {
session_id,
run_id,
tool_call_id,
tool_name,
metadata,
trace_context: None,
policy: None,
}
}
#[must_use]
pub const fn with_trace_context(mut self, trace_context: &'a TraceContext) -> Self {
self.trace_context = Some(trace_context);
self
}
#[must_use]
pub fn with_policy(mut self, policy: Value) -> Self {
self.policy = Some(policy);
self
}
fn control_flow(&self) -> Option<&str> {
self.metadata.get("control_flow").and_then(Value::as_str)
}
fn approval_id(&self) -> String {
format!("approval_{}_{}", self.run_id.as_str(), self.tool_call_id)
}
fn deferred_id(&self) -> String {
format!("deferred_{}_{}", self.run_id.as_str(), self.tool_call_id)
}
}
#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum ApprovalStatus {
#[default]
Pending,
Approved,
Denied,
Expired,
Cancelled,
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct ApprovalDecision {
pub status: ApprovalStatus,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub decided_by: Option<String>,
pub decided_at: DateTime<Utc>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub reason: Option<String>,
#[serde(default, skip_serializing_if = "Metadata::is_empty")]
pub metadata: Metadata,
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct ApprovalRecord {
pub approval_id: String,
pub session_id: SessionId,
pub run_id: RunId,
pub action_id: String,
pub action_name: String,
#[serde(default, skip_serializing_if = "Value::is_null")]
pub request: Value,
#[serde(default)]
pub status: ApprovalStatus,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub decision: Option<ApprovalDecision>,
pub created_at: DateTime<Utc>,
pub updated_at: DateTime<Utc>,
#[serde(default, skip_serializing_if = "TraceContext::is_empty")]
pub trace_context: TraceContext,
#[serde(default, skip_serializing_if = "Metadata::is_empty")]
pub metadata: Metadata,
}
impl ApprovalRecord {
#[must_use]
pub fn new(
approval_id: impl Into<String>,
session_id: SessionId,
run_id: RunId,
action_id: impl Into<String>,
action_name: impl Into<String>,
) -> Self {
let now = Utc::now();
Self {
approval_id: approval_id.into(),
session_id,
run_id,
action_id: action_id.into(),
action_name: action_name.into(),
request: Value::Null,
status: ApprovalStatus::Pending,
decision: None,
created_at: now,
updated_at: now,
trace_context: TraceContext::default(),
metadata: Metadata::default(),
}
}
#[must_use]
pub fn from_tool_return(input: &ToolReturnRecordInput<'_>) -> Option<Self> {
if input.control_flow() != Some("approval_required") {
return None;
}
let mut record = Self::new(
input.approval_id(),
input.session_id.clone(),
input.run_id.clone(),
input.tool_call_id,
input.tool_name,
);
record.request = input
.metadata
.get("approval")
.cloned()
.unwrap_or(Value::Null);
if let Some(trace_context) = input.trace_context {
record.trace_context = trace_context.clone();
}
if let Some(policy) = &input.policy {
record.metadata.insert("policy".to_string(), policy.clone());
}
Some(record)
}
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
#[serde(tag = "decision", rename_all = "snake_case")]
pub enum ToolApprovalDecision {
Approved {
#[serde(default, skip_serializing_if = "Option::is_none")]
decided_by: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
reason: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
override_arguments: Option<Value>,
#[serde(default, skip_serializing_if = "Metadata::is_empty")]
metadata: Metadata,
},
Denied {
#[serde(default, skip_serializing_if = "Option::is_none")]
decided_by: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
reason: Option<String>,
#[serde(default, skip_serializing_if = "Metadata::is_empty")]
metadata: Metadata,
},
}
impl ToolApprovalDecision {
#[must_use]
pub fn approved() -> Self {
Self::Approved {
decided_by: None,
reason: None,
override_arguments: None,
metadata: Metadata::default(),
}
}
#[must_use]
pub fn denied(reason: impl Into<String>) -> Self {
Self::Denied {
decided_by: None,
reason: Some(reason.into()),
metadata: Metadata::default(),
}
}
#[must_use]
pub fn with_override_arguments(self, arguments: Value) -> Self {
match self {
Self::Approved {
decided_by,
reason,
metadata,
..
} => Self::Approved {
decided_by,
reason,
override_arguments: Some(arguments),
metadata,
},
denied @ Self::Denied { .. } => denied,
}
}
#[must_use]
pub fn into_approval_decision(self) -> ApprovalDecision {
let decided_at = Utc::now();
match self {
Self::Approved {
decided_by,
reason,
override_arguments,
mut metadata,
} => {
if let Some(arguments) = override_arguments {
metadata.insert("override_arguments".to_string(), arguments);
}
ApprovalDecision {
status: ApprovalStatus::Approved,
decided_by,
decided_at,
reason,
metadata,
}
}
Self::Denied {
decided_by,
reason,
metadata,
} => ApprovalDecision {
status: ApprovalStatus::Denied,
decided_by,
decided_at,
reason,
metadata,
},
}
}
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct DeferredToolRecord {
pub deferred_id: String,
pub session_id: SessionId,
pub run_id: RunId,
pub tool_call_id: String,
pub tool_name: String,
#[serde(default, skip_serializing_if = "Value::is_null")]
pub request: Value,
pub status: ExecutionStatus,
#[serde(default, skip_serializing_if = "Value::is_null")]
pub response: Value,
pub created_at: DateTime<Utc>,
pub updated_at: DateTime<Utc>,
#[serde(default, skip_serializing_if = "TraceContext::is_empty")]
pub trace_context: TraceContext,
#[serde(default, skip_serializing_if = "Metadata::is_empty")]
pub metadata: Metadata,
}
impl DeferredToolRecord {
#[must_use]
pub fn new(
deferred_id: impl Into<String>,
session_id: SessionId,
run_id: RunId,
tool_call_id: impl Into<String>,
tool_name: impl Into<String>,
) -> Self {
let now = Utc::now();
Self {
deferred_id: deferred_id.into(),
session_id,
run_id,
tool_call_id: tool_call_id.into(),
tool_name: tool_name.into(),
request: Value::Null,
status: ExecutionStatus::Pending,
response: Value::Null,
created_at: now,
updated_at: now,
trace_context: TraceContext::default(),
metadata: Metadata::default(),
}
}
#[must_use]
pub fn from_tool_return(input: &ToolReturnRecordInput<'_>) -> Option<Self> {
if input.control_flow() != Some("call_deferred") {
return None;
}
let mut record = Self::new(
input.deferred_id(),
input.session_id.clone(),
input.run_id.clone(),
input.tool_call_id,
input.tool_name,
);
record.request = input
.metadata
.get("deferred")
.cloned()
.unwrap_or(Value::Null);
record.status = ExecutionStatus::Waiting;
if let Some(trace_context) = input.trace_context {
record.trace_context = trace_context.clone();
}
if let Some(policy) = &input.policy {
record.metadata.insert("policy".to_string(), policy.clone());
}
Some(record)
}
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct DeferredToolRequest {
pub deferred_id: String,
pub session_id: SessionId,
pub run_id: RunId,
pub tool_call_id: String,
pub tool_name: String,
#[serde(default, skip_serializing_if = "Value::is_null")]
pub arguments: Value,
#[serde(default, skip_serializing_if = "TraceContext::is_empty")]
pub trace_context: TraceContext,
#[serde(default, skip_serializing_if = "Metadata::is_empty")]
pub metadata: Metadata,
}
impl DeferredToolRequest {
#[must_use]
pub fn from_record(record: &DeferredToolRecord) -> Self {
Self {
deferred_id: record.deferred_id.clone(),
session_id: record.session_id.clone(),
run_id: record.run_id.clone(),
tool_call_id: record.tool_call_id.clone(),
tool_name: record.tool_name.clone(),
arguments: record.request.clone(),
trace_context: record.trace_context.clone(),
metadata: record.metadata.clone(),
}
}
#[must_use]
pub fn into_record(self) -> DeferredToolRecord {
let mut record = DeferredToolRecord::new(
self.deferred_id,
self.session_id,
self.run_id,
self.tool_call_id,
self.tool_name,
);
record.request = self.arguments;
record.trace_context = self.trace_context;
record.metadata = self.metadata;
record
}
}
#[derive(Clone, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
pub struct DeferredToolRequests {
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub requests: Vec<DeferredToolRequest>,
}
impl DeferredToolRequests {
#[must_use]
pub fn from_records(records: &[DeferredToolRecord]) -> Self {
Self {
requests: records
.iter()
.map(DeferredToolRequest::from_record)
.collect(),
}
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.requests.is_empty()
}
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct DeferredToolResult {
pub deferred_id: String,
pub status: ExecutionStatus,
#[serde(default, skip_serializing_if = "Value::is_null")]
pub response: Value,
#[serde(default, skip_serializing_if = "Metadata::is_empty")]
pub metadata: Metadata,
}
impl DeferredToolResult {
#[must_use]
pub fn completed(deferred_id: impl Into<String>, response: Value) -> Self {
Self {
deferred_id: deferred_id.into(),
status: ExecutionStatus::Completed,
response,
metadata: Metadata::default(),
}
}
#[must_use]
pub fn failed(deferred_id: impl Into<String>, response: Value) -> Self {
Self {
deferred_id: deferred_id.into(),
status: ExecutionStatus::Failed,
response,
metadata: Metadata::default(),
}
}
#[must_use]
pub fn cancelled(deferred_id: impl Into<String>, response: Value) -> Self {
Self {
deferred_id: deferred_id.into(),
status: ExecutionStatus::Cancelled,
response,
metadata: Metadata::default(),
}
}
#[must_use]
pub fn with_metadata(mut self, metadata: Metadata) -> Self {
self.metadata = metadata;
self
}
pub fn apply_to_record(self, record: &mut DeferredToolRecord) {
record.status = self.status;
record.response = self.response;
record.updated_at = Utc::now();
record.metadata.extend(self.metadata);
}
}
#[derive(Clone, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
pub struct DeferredToolResults {
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub results: Vec<DeferredToolResult>,
}
impl DeferredToolResults {
#[must_use]
pub fn new(results: impl IntoIterator<Item = DeferredToolResult>) -> Self {
Self {
results: results.into_iter().collect(),
}
}
#[must_use]
pub fn is_empty(&self) -> bool {
self.results.is_empty()
}
}