use chrono::{DateTime, Utc};
use serde::{Deserialize, Deserializer, Serialize, Serializer, de::Error as _};
use serde_json::Value;
use starweaver_context::ResumableState;
use starweaver_core::{
CheckpointId, ConversationId, Metadata, RunId, RunLifecycle, SessionId, TaskId, TraceContext,
};
use starweaver_stream::{ReplayCursor, ReplayCursorFamily, ReplayScope};
use crate::{input::InputPart, management::SessionDeletionFence};
#[derive(Clone, Copy, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum SessionStatus {
#[default]
Active,
Archived,
Failed,
Deleted,
}
#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
pub struct DurableRunStatus(Option<RunLifecycle>);
pub type RunStatus = DurableRunStatus;
#[allow(non_upper_case_globals)]
impl DurableRunStatus {
pub const Queued: Self = Self(None);
pub const Starting: Self = Self(Some(RunLifecycle::Starting));
pub const Running: Self = Self(Some(RunLifecycle::Running));
pub const Waiting: Self = Self(Some(RunLifecycle::Waiting));
pub const Completed: Self = Self(Some(RunLifecycle::Completed));
pub const Failed: Self = Self(Some(RunLifecycle::Failed));
pub const Cancelled: Self = Self(Some(RunLifecycle::Cancelled));
#[must_use]
pub const fn lifecycle(self) -> Option<RunLifecycle> {
self.0
}
#[must_use]
pub const fn as_str(self) -> &'static str {
match self.0 {
None => "queued",
Some(lifecycle) => lifecycle.as_str(),
}
}
#[must_use]
pub const fn is_active(self) -> bool {
matches!(
self.0,
None | Some(RunLifecycle::Starting | RunLifecycle::Running | RunLifecycle::Waiting)
)
}
#[must_use]
pub const fn is_terminal(self) -> bool {
match self.0 {
Some(lifecycle) => lifecycle.is_terminal(),
None => false,
}
}
}
impl From<RunLifecycle> for DurableRunStatus {
fn from(value: RunLifecycle) -> Self {
Self(Some(value))
}
}
impl TryFrom<DurableRunStatus> for RunLifecycle {
type Error = QueuedRunStatus;
fn try_from(value: DurableRunStatus) -> Result<Self, Self::Error> {
value.0.ok_or(QueuedRunStatus)
}
}
impl Serialize for DurableRunStatus {
fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
where
S: Serializer,
{
serializer.serialize_str(self.as_str())
}
}
impl<'de> Deserialize<'de> for DurableRunStatus {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
match String::deserialize(deserializer)?.as_str() {
"queued" => Ok(Self::Queued),
"starting" => Ok(Self::Starting),
"running" => Ok(Self::Running),
"waiting" => Ok(Self::Waiting),
"completed" => Ok(Self::Completed),
"failed" => Ok(Self::Failed),
"cancelled" => Ok(Self::Cancelled),
other => Err(D::Error::custom(format!("unknown run status: {other}"))),
}
}
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct RunTerminalError {
pub code: String,
pub message: String,
}
impl RunTerminalError {
#[must_use]
pub fn new(code: impl Into<String>, message: impl Into<String>) -> Self {
Self {
code: code.into(),
message: message.into(),
}
}
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct RunTerminalProjection {
pub status: RunStatus,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_preview: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub error: Option<RunTerminalError>,
}
impl RunTerminalProjection {
pub fn try_new(
status: RunStatus,
output_preview: Option<String>,
error: Option<RunTerminalError>,
) -> Result<Self, RunTerminalProjectionError> {
let projection = Self {
status,
output_preview,
error,
};
projection.validate()?;
Ok(projection)
}
#[must_use]
pub const fn completed(output_preview: Option<String>) -> Self {
Self {
status: RunStatus::Completed,
output_preview,
error: None,
}
}
#[must_use]
pub const fn failed(error: RunTerminalError) -> Self {
Self {
status: RunStatus::Failed,
output_preview: None,
error: Some(error),
}
}
#[must_use]
pub const fn cancelled(error: Option<RunTerminalError>) -> Self {
Self {
status: RunStatus::Cancelled,
output_preview: None,
error,
}
}
pub fn validate(&self) -> Result<(), RunTerminalProjectionError> {
if !self.status.is_terminal() {
return Err(RunTerminalProjectionError::NonTerminalStatus(self.status));
}
if self.status == RunStatus::Failed && self.error.is_none() {
return Err(RunTerminalProjectionError::MissingFailureDiagnostic);
}
if self.status == RunStatus::Completed && self.error.is_some() {
return Err(RunTerminalProjectionError::UnexpectedSuccessDiagnostic);
}
if let Some(error) = self.error.as_ref() {
if error.code.is_empty() {
return Err(RunTerminalProjectionError::EmptyDiagnosticCode);
}
if error.message.is_empty() {
return Err(RunTerminalProjectionError::EmptyDiagnosticMessage);
}
}
Ok(())
}
#[must_use]
pub fn matches(&self, run: &RunRecord) -> bool {
(run.status, &run.output_preview, &run.terminal_error)
== (self.status, &self.output_preview, &self.error)
}
pub fn apply_to(&self, run: &mut RunRecord) {
run.status = self.status;
run.output_preview.clone_from(&self.output_preview);
run.terminal_error.clone_from(&self.error);
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum RunTerminalProjectionError {
NonTerminalStatus(RunStatus),
MissingFailureDiagnostic,
UnexpectedSuccessDiagnostic,
UnexpectedNonTerminalDiagnostic(RunStatus),
EmptyDiagnosticCode,
EmptyDiagnosticMessage,
}
impl std::fmt::Display for RunTerminalProjectionError {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::NonTerminalStatus(status) => {
write!(formatter, "run status {} is not terminal", status.as_str())
}
Self::MissingFailureDiagnostic => {
formatter.write_str("failed run requires a terminal diagnostic")
}
Self::UnexpectedSuccessDiagnostic => {
formatter.write_str("completed run cannot carry a terminal diagnostic")
}
Self::UnexpectedNonTerminalDiagnostic(status) => write!(
formatter,
"non-terminal run status {} cannot carry a terminal diagnostic",
status.as_str()
),
Self::EmptyDiagnosticCode => formatter.write_str("terminal diagnostic code is empty"),
Self::EmptyDiagnosticMessage => {
formatter.write_str("terminal diagnostic message is empty")
}
}
}
}
impl std::error::Error for RunTerminalProjectionError {}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub struct QueuedRunStatus;
impl std::fmt::Display for QueuedRunStatus {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
formatter.write_str("queued run has no runtime lifecycle")
}
}
impl std::error::Error for QueuedRunStatus {}
#[derive(Clone, Copy, Debug, Deserialize, Eq, PartialEq, Serialize)]
#[serde(rename_all = "snake_case")]
pub enum ExecutionStatus {
Pending,
Running,
Waiting,
Completed,
Failed,
Cancelled,
}
#[derive(Clone, Debug, Default, Deserialize, Eq, PartialEq, Serialize)]
pub struct EnvironmentStateRef {
pub provider: String,
pub reference: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub revision: Option<String>,
#[serde(default, skip_serializing_if = "Metadata::is_empty")]
pub metadata: Metadata,
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct CheckpointRef {
pub checkpoint_id: CheckpointId,
pub run_id: RunId,
pub sequence: usize,
pub node: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub storage_ref: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub stream_cursor: Option<usize>,
pub created_at: DateTime<Utc>,
#[serde(default, skip_serializing_if = "Metadata::is_empty")]
pub metadata: Metadata,
}
#[derive(Clone, Debug, Eq, PartialEq, Serialize)]
pub struct StreamCursorRef {
pub position: ReplayCursor,
pub created_at: DateTime<Utc>,
#[serde(default, skip_serializing_if = "Metadata::is_empty")]
pub metadata: Metadata,
}
impl StreamCursorRef {
#[must_use]
pub fn new(position: ReplayCursor) -> Self {
Self {
position,
created_at: Utc::now(),
metadata: Metadata::default(),
}
}
#[must_use]
pub const fn family(&self) -> ReplayCursorFamily {
self.position.family
}
#[must_use]
pub const fn scope(&self) -> &ReplayScope {
&self.position.scope
}
#[must_use]
pub const fn sequence(&self) -> usize {
self.position.sequence
}
#[must_use]
pub fn same_stream(&self, other: &Self) -> bool {
self.family() == other.family() && self.scope() == other.scope()
}
pub fn validate_for_run(&self, run_id: &RunId) -> Result<(), StreamCursorRefError> {
let expected = ReplayScope::run(run_id.as_str());
if self.scope() != &expected {
return Err(StreamCursorRefError::WrongScope {
expected: expected.as_str().to_string(),
actual: self.scope().as_str().to_string(),
});
}
Ok(())
}
pub fn validate_progression(&self, current: &Self) -> Result<(), StreamCursorRefError> {
if self.same_stream(current) && self.sequence() < current.sequence() {
return Err(StreamCursorRefError::SequenceRegression {
family: self.family(),
scope: self.scope().as_str().to_string(),
current: current.sequence(),
proposed: self.sequence(),
});
}
Ok(())
}
}
#[derive(Clone, Debug, Eq, PartialEq)]
pub enum StreamCursorRefError {
WrongScope {
expected: String,
actual: String,
},
SequenceRegression {
family: ReplayCursorFamily,
scope: String,
current: usize,
proposed: usize,
},
}
impl std::fmt::Display for StreamCursorRefError {
fn fmt(&self, formatter: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::WrongScope { expected, actual } => {
write!(
formatter,
"expected cursor scope {expected}, received {actual}"
)
}
Self::SequenceRegression {
family,
scope,
current,
proposed,
} => write!(
formatter,
"{} cursor for {scope} regressed from {current} to {proposed}",
family.as_str()
),
}
}
}
impl std::error::Error for StreamCursorRefError {}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct CurrentStreamCursorRefWire {
position: ReplayCursor,
created_at: DateTime<Utc>,
#[serde(default)]
metadata: Metadata,
}
#[derive(Deserialize)]
#[serde(deny_unknown_fields)]
struct LegacyStreamCursorRefWire {
family: String,
scope: String,
sequence: usize,
#[serde(default)]
cursor: Option<String>,
created_at: DateTime<Utc>,
#[serde(default)]
metadata: Metadata,
}
#[derive(Deserialize)]
#[serde(untagged)]
enum StreamCursorRefWire {
Current(CurrentStreamCursorRefWire),
Legacy(LegacyStreamCursorRefWire),
}
impl<'de> Deserialize<'de> for StreamCursorRef {
fn deserialize<D>(deserializer: D) -> Result<Self, D::Error>
where
D: Deserializer<'de>,
{
match StreamCursorRefWire::deserialize(deserializer)? {
StreamCursorRefWire::Current(current) => Ok(Self {
position: current.position,
created_at: current.created_at,
metadata: current.metadata,
}),
StreamCursorRefWire::Legacy(legacy) => {
let LegacyStreamCursorRefWire {
family,
scope,
sequence,
cursor,
created_at,
metadata,
} = legacy;
let family = match family.as_str() {
"raw_runtime" => ReplayCursorFamily::RawRuntime,
"display" => ReplayCursorFamily::Display,
"replay_event" => ReplayCursorFamily::ReplayEvent,
other => {
return Err(D::Error::custom(format!(
"unknown stream cursor family: {other}"
)));
}
};
let mut position =
ReplayCursor::for_family(family, ReplayScope::from_string(scope), sequence);
position.backend_cursor = cursor;
Ok(Self {
position,
created_at,
metadata,
})
}
}
}
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct SessionRecord {
pub session_id: SessionId,
#[serde(default = "default_session_namespace")]
pub namespace_id: String,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub owner_id: Option<String>,
#[serde(default = "initial_session_revision")]
pub revision: u64,
#[serde(default)]
pub deletion_fence: SessionDeletionFence,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub title: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub workspace: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub profile: Option<String>,
#[serde(default)]
pub status: SessionStatus,
#[serde(default)]
pub state: ResumableState,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub environment_state: Option<EnvironmentStateRef>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub stream_cursors: Vec<StreamCursorRef>,
#[serde(default, skip_serializing_if = "TraceContext::is_empty")]
pub trace_context: TraceContext,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub parent_session_id: Option<SessionId>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub head_run_id: Option<RunId>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub head_success_run_id: Option<RunId>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub active_run_id: Option<RunId>,
pub created_at: DateTime<Utc>,
pub updated_at: DateTime<Utc>,
#[serde(default, skip_serializing_if = "Metadata::is_empty")]
pub metadata: Metadata,
}
fn default_session_namespace() -> String {
crate::LOCAL_SESSION_NAMESPACE.to_string()
}
const fn initial_session_revision() -> u64 {
1
}
const fn initial_run_revision() -> u64 {
1
}
impl starweaver_core::VersionedRecord for SessionRecord {
const SCHEMA: &'static str = "starweaver.session.session_record";
const ALLOW_BARE_V0: bool = true;
}
impl SessionRecord {
#[must_use]
pub fn new(session_id: SessionId) -> Self {
let now = Utc::now();
Self {
session_id,
namespace_id: default_session_namespace(),
owner_id: None,
revision: initial_session_revision(),
deletion_fence: SessionDeletionFence::Stable,
title: None,
workspace: None,
profile: None,
status: SessionStatus::Active,
state: ResumableState::default(),
environment_state: None,
stream_cursors: Vec::new(),
trace_context: TraceContext::default(),
parent_session_id: None,
head_run_id: None,
head_success_run_id: None,
active_run_id: None,
created_at: now,
updated_at: now,
metadata: Metadata::default(),
}
}
}
#[derive(Clone, Debug, Deserialize, Eq, PartialEq, Serialize)]
pub struct RunRecord {
pub session_id: SessionId,
pub run_id: RunId,
#[serde(default = "initial_run_revision")]
pub revision: u64,
pub conversation_id: ConversationId,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub input: Vec<InputPart>,
#[serde(default)]
pub status: RunStatus,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub output_preview: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub terminal_error: Option<RunTerminalError>,
#[serde(default, skip_serializing_if = "Value::is_null")]
pub structured_output: Value,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub latest_checkpoint: Option<CheckpointRef>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub environment_state: Option<EnvironmentStateRef>,
#[serde(default, skip_serializing_if = "Vec::is_empty")]
pub stream_cursors: Vec<StreamCursorRef>,
#[serde(default, skip_serializing_if = "TraceContext::is_empty")]
pub trace_context: TraceContext,
#[serde(default)]
pub sequence_no: usize,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub restore_from_run_id: Option<RunId>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub parent_run_id: Option<RunId>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub parent_task_id: Option<TaskId>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub trigger_type: Option<String>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub profile: Option<String>,
pub created_at: DateTime<Utc>,
pub updated_at: DateTime<Utc>,
#[serde(default, skip_serializing_if = "Metadata::is_empty")]
pub metadata: Metadata,
}
impl starweaver_core::VersionedRecord for RunRecord {
const SCHEMA: &'static str = "starweaver.session.run_record";
const ALLOW_BARE_V0: bool = true;
}
impl RunRecord {
#[must_use]
pub fn terminal_projection(&self) -> Option<RunTerminalProjection> {
self.status.is_terminal().then(|| RunTerminalProjection {
status: self.status,
output_preview: self.output_preview.clone(),
error: self.terminal_error.clone(),
})
}
pub fn validate_new_write(&self) -> Result<(), RunTerminalProjectionError> {
self.terminal_projection().map_or_else(
|| {
if self.terminal_error.is_some() {
Err(RunTerminalProjectionError::UnexpectedNonTerminalDiagnostic(
self.status,
))
} else {
Ok(())
}
},
|terminal| terminal.validate(),
)
}
pub fn normalize_for_admission(&mut self) {
self.status = RunStatus::Queued;
self.output_preview = None;
self.terminal_error = None;
}
pub fn apply_legacy_status_update(
&mut self,
status: RunStatus,
output_preview: Option<String>,
) {
self.status = status;
match status {
RunStatus::Failed => {
self.output_preview = None;
self.terminal_error = Some(RunTerminalError::new(
"legacy_status_update_failed",
"run failed",
));
}
RunStatus::Cancelled => {
self.output_preview = None;
self.terminal_error = Some(RunTerminalError::new(
"legacy_status_update_cancelled",
"run cancelled",
));
}
_ => {
self.output_preview = output_preview;
self.terminal_error = None;
}
}
}
#[must_use]
pub fn new(session_id: SessionId, run_id: RunId, conversation_id: ConversationId) -> Self {
let now = Utc::now();
Self {
session_id,
run_id,
revision: initial_run_revision(),
conversation_id,
input: Vec::new(),
status: RunStatus::Queued,
output_preview: None,
terminal_error: None,
structured_output: Value::Null,
latest_checkpoint: None,
environment_state: None,
stream_cursors: Vec::new(),
trace_context: TraceContext::default(),
sequence_no: 0,
restore_from_run_id: None,
parent_run_id: None,
parent_task_id: None,
trigger_type: None,
profile: None,
created_at: now,
updated_at: now,
metadata: Metadata::default(),
}
}
}