use std::fmt;
use std::future::Future;
use std::path::PathBuf;
use std::pin::Pin;
use std::sync::Arc;
use ag_protocol::{AgentResponse, ProtocolRequestProfile, TurnPrompt};
use tokio::sync::mpsc;
use crate::model::agent::ReasoningLevel;
pub type AgentFuture<T> = Pin<Box<dyn Future<Output = T> + Send>>;
pub trait LiveTranscript: fmt::Debug + Send + Sync {
fn replay_text(&self) -> Option<String>;
}
#[derive(Debug, Clone, PartialEq, Eq)]
pub enum AgentRequestKind {
SessionStart,
SessionResume,
UtilityPrompt,
AccountRead,
}
impl AgentRequestKind {
#[must_use]
pub fn protocol_profile(&self) -> ProtocolRequestProfile {
match self {
Self::SessionStart | Self::SessionResume => ProtocolRequestProfile::SessionTurn,
Self::UtilityPrompt | Self::AccountRead => ProtocolRequestProfile::UtilityPrompt,
}
}
#[must_use]
pub fn is_resume(&self) -> bool {
matches!(self, Self::SessionResume)
}
}
#[derive(Clone, Debug)]
pub struct TurnContinuation {
kind: TurnContinuationKind,
}
impl TurnContinuation {
#[must_use]
pub fn fresh() -> Self {
Self {
kind: TurnContinuationKind::Fresh,
}
}
#[must_use]
pub fn replaying(replay_transcript: String) -> Self {
Self {
kind: TurnContinuationKind::Replay { replay_transcript },
}
}
#[must_use]
pub fn provider(
live_transcript: Option<Arc<dyn LiveTranscript>>,
persisted_instruction_conversation_id: Option<String>,
provider_conversation_id: Option<String>,
replay_transcript: Option<String>,
) -> Self {
Self {
kind: TurnContinuationKind::Provider {
live_transcript,
persisted_instruction_conversation_id,
provider_conversation_id,
replay_transcript,
},
}
}
#[must_use]
pub fn replay_transcript(&self) -> Option<&str> {
match &self.kind {
TurnContinuationKind::Fresh => None,
TurnContinuationKind::Provider {
replay_transcript, ..
} => replay_transcript.as_deref(),
TurnContinuationKind::Replay { replay_transcript } => Some(replay_transcript.as_str()),
}
}
#[must_use]
pub fn provider_conversation_id(&self) -> Option<&str> {
match &self.kind {
TurnContinuationKind::Provider {
provider_conversation_id,
..
} => provider_conversation_id.as_deref(),
TurnContinuationKind::Fresh | TurnContinuationKind::Replay { .. } => None,
}
}
#[must_use]
pub fn persisted_instruction_conversation_id(&self) -> Option<&str> {
match &self.kind {
TurnContinuationKind::Provider {
persisted_instruction_conversation_id,
..
} => persisted_instruction_conversation_id.as_deref(),
TurnContinuationKind::Fresh | TurnContinuationKind::Replay { .. } => None,
}
}
pub(crate) fn into_parts(self) -> TurnContinuationParts {
match self.kind {
TurnContinuationKind::Fresh => TurnContinuationParts::default(),
TurnContinuationKind::Replay { replay_transcript } => TurnContinuationParts {
replay_transcript: Some(replay_transcript),
..TurnContinuationParts::default()
},
TurnContinuationKind::Provider {
live_transcript,
persisted_instruction_conversation_id,
provider_conversation_id,
replay_transcript,
} => TurnContinuationParts {
live_transcript,
persisted_instruction_conversation_id,
provider_conversation_id,
replay_transcript,
},
}
}
}
#[derive(Clone, Debug)]
enum TurnContinuationKind {
Fresh,
Provider {
live_transcript: Option<Arc<dyn LiveTranscript>>,
persisted_instruction_conversation_id: Option<String>,
provider_conversation_id: Option<String>,
replay_transcript: Option<String>,
},
Replay {
replay_transcript: String,
},
}
#[derive(Default)]
pub(crate) struct TurnContinuationParts {
pub(crate) live_transcript: Option<Arc<dyn LiveTranscript>>,
pub(crate) persisted_instruction_conversation_id: Option<String>,
pub(crate) provider_conversation_id: Option<String>,
pub(crate) replay_transcript: Option<String>,
}
#[derive(Debug, Clone)]
pub struct TurnRequest {
pub continuation: TurnContinuation,
pub folder: PathBuf,
pub main_checkout_root: Option<PathBuf>,
pub model: String,
pub prompt: TurnPrompt,
pub reasoning_level: ReasoningLevel,
pub request_kind: AgentRequestKind,
}
#[derive(Clone, Debug, PartialEq)]
pub enum TurnEvent {
ThoughtDelta(String),
Completed {
context_reset: bool,
input_tokens: u64,
output_tokens: u64,
},
Failed(String),
PidUpdate(Option<u32>),
}
#[derive(Debug)]
pub struct TurnResult {
pub assistant_message: AgentResponse,
pub context_reset: bool,
pub input_tokens: u64,
pub output_tokens: u64,
pub provider_conversation_id: Option<String>,
}
pub struct SessionRef {
pub session_id: String,
}
pub struct StartSessionRequest {
pub folder: PathBuf,
pub session_id: String,
}
#[derive(Debug, thiserror::Error)]
pub enum AgentError {
#[error(transparent)]
AppServer(#[from] crate::app_server::AppServerError),
#[error("{0}")]
Backend(String),
#[error("{0}")]
InterruptedByUser(String),
#[error("{0}")]
Io(String),
}
#[cfg_attr(any(test, feature = "test-utils"), mockall::automock)]
pub trait AgentChannel: Send + Sync {
fn start_session(
&self,
req: StartSessionRequest,
) -> AgentFuture<Result<SessionRef, AgentError>>;
fn run_turn(
&self,
session_id: String,
req: TurnRequest,
events: mpsc::UnboundedSender<TurnEvent>,
) -> AgentFuture<Result<TurnResult, AgentError>>;
fn shutdown_session(&self, session_id: String) -> AgentFuture<Result<(), AgentError>>;
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn test_turn_continuation_fresh_has_no_context() {
let continuation = TurnContinuation::fresh();
assert_eq!(continuation.replay_transcript(), None);
assert_eq!(continuation.provider_conversation_id(), None);
assert_eq!(continuation.persisted_instruction_conversation_id(), None);
}
#[test]
fn test_turn_continuation_replaying_exposes_transcript_only() {
let continuation = TurnContinuation::replaying("prior turn".to_string());
let parts = continuation.clone().into_parts();
assert_eq!(continuation.replay_transcript(), Some("prior turn"));
assert_eq!(continuation.provider_conversation_id(), None);
assert!(parts.live_transcript.is_none());
assert_eq!(parts.persisted_instruction_conversation_id, None);
assert_eq!(parts.provider_conversation_id, None);
assert_eq!(parts.replay_transcript.as_deref(), Some("prior turn"));
}
#[test]
fn test_turn_continuation_provider_exposes_persisted_context() {
let continuation = TurnContinuation::provider(
None,
Some("instruction-1".to_string()),
Some("thread-1".to_string()),
Some("prior turn".to_string()),
);
assert_eq!(continuation.replay_transcript(), Some("prior turn"));
assert_eq!(continuation.provider_conversation_id(), Some("thread-1"));
assert_eq!(
continuation.persisted_instruction_conversation_id(),
Some("instruction-1")
);
}
#[test]
fn test_agent_request_kind_session_variants_use_session_protocol_profile() {
let start = AgentRequestKind::SessionStart;
let resume = AgentRequestKind::SessionResume;
let start_profile = start.protocol_profile();
let resume_profile = resume.protocol_profile();
assert_eq!(start_profile, ProtocolRequestProfile::SessionTurn);
assert_eq!(resume_profile, ProtocolRequestProfile::SessionTurn);
}
#[test]
fn test_agent_request_kind_utility_prompt_uses_utility_protocol_profile() {
let request_kind = AgentRequestKind::UtilityPrompt;
let protocol_profile = request_kind.protocol_profile();
assert_eq!(protocol_profile, ProtocolRequestProfile::UtilityPrompt);
}
#[test]
fn test_agent_request_kind_account_read_uses_utility_protocol_profile() {
let request_kind = AgentRequestKind::AccountRead;
let protocol_profile = request_kind.protocol_profile();
assert_eq!(protocol_profile, ProtocolRequestProfile::UtilityPrompt);
}
}