1use std::collections::BTreeMap;
4use std::sync::Arc;
5use std::sync::Mutex;
6use std::sync::Weak;
7use std::time::SystemTime;
8use std::time::UNIX_EPOCH;
9
10use serde_json::Value;
11use tokio::sync::{mpsc, oneshot};
12use uuid::Uuid;
13
14use crate::Error;
15use crate::Result;
16use crate::backend::checkpoint::CHECKPOINT_VERSION;
17use crate::backend::checkpoint::Checkpoint;
18use crate::backend::checkpoint::CheckpointStore;
19use crate::backend::checkpoint::ExecutionOutcome;
20use crate::backend::checkpoint::ExecutionRecord;
21use crate::backend::checkpoint::JournalEvent;
22use crate::backend::model::ModelRouter;
23use crate::backend::sandbox::Sandbox;
24use crate::backend::session_files::session_file_limits;
25use crate::middleware::FrontendExtensions;
26use crate::middleware::MiddlewareCommandContext;
27use crate::middleware::MiddlewareStack;
28use crate::middleware::RuntimeContext;
29use crate::middleware::tools::Catalog;
30use crate::protocol::Event;
31use crate::protocol::EventMsg;
32use crate::protocol::ModelChangedEvent;
33use crate::protocol::ModelChoice;
34use crate::protocol::ModelInfo;
35use crate::protocol::Op;
36use crate::protocol::SessionConfiguredEvent;
37use crate::protocol::SessionContext;
38use crate::protocol::SessionResumeRequestedEvent;
39use crate::protocol::Submission;
40use crate::protocol::TokenUsage;
41use crate::protocol::WarningEvent;
42
43mod approval;
44mod input;
45mod recorder;
46mod startup;
47mod tool_step;
48mod turn;
49
50pub use self::startup::create_agent;
51
52use self::recorder::EventRecorder;
53
54const SUBMISSION_QUEUE_CAPACITY: usize = 64;
55const EVENT_QUEUE_CAPACITY: usize = 256;
56const MAX_OPERATION_BYTES: usize = 256;
57const DEFAULT_INITIAL_REPLAY_BATCHES: usize = 100;
58
59type UsageObserver =
60 Arc<dyn for<'a> Fn(&'a str, &'a TokenUsage) -> crate::BoxFuture<'a, Result<()>> + Send + Sync>;
61
62pub const DEFAULT_MAX_MODEL_STEPS: usize = 2042;
64
65pub const TURN_INTERRUPTED_REASON: &str = "interrupted";
67pub const TURN_RESTARTED_REASON: &str = "interrupted by restart";
69pub const FRONTEND_DISCONNECTED_REASON: &str = "frontend disconnected";
71
72#[derive(Debug, Clone, Default, PartialEq, Eq)]
74pub enum AgentRole {
75 #[default]
76 Main,
78 Subagent {
80 parent_session_id: String,
82 parent_turn_id: String,
84 },
85}
86
87#[derive(Clone)]
89pub struct AgentConfig {
90 model: Arc<ModelRouter>,
91 provider: String,
92 sandbox: Arc<Sandbox>,
93 checkpoints: Arc<dyn CheckpointStore>,
94 middleware: MiddlewareStack,
95 system_prompt: String,
96 session_id: String,
97 context_window: i64,
98 default_context_window: i64,
99 session_context: SessionContext,
100 metadata: BTreeMap<String, Value>,
101 catalog_visible: bool,
102 usage_observer: Option<UsageObserver>,
103 metadata_configured: bool,
104 model_route_configured: bool,
105 initial_replay_batches: usize,
106 max_model_steps: usize,
107 token_estimate: crate::middleware::TokenEstimate,
108 role: AgentRole,
109}
110
111impl AgentConfig {
112 pub fn new(
114 model: Arc<ModelRouter>,
115 sandbox: Arc<Sandbox>,
116 checkpoints: Arc<dyn CheckpointStore>,
117 middleware: MiddlewareStack,
118 system_prompt: impl Into<String>,
119 ) -> Self {
120 let provider = model.default_provider().to_string();
121 let session_id = Uuid::new_v4().to_string();
122 Self {
123 model,
124 provider,
125 sandbox,
126 checkpoints,
127 middleware,
128 system_prompt: system_prompt.into(),
129 session_id,
130 context_window: 272_000,
131 default_context_window: 272_000,
132 session_context: SessionContext::default(),
133 metadata: BTreeMap::new(),
134 catalog_visible: true,
135 usage_observer: None,
136 metadata_configured: false,
137 model_route_configured: false,
138 initial_replay_batches: DEFAULT_INITIAL_REPLAY_BATCHES,
139 max_model_steps: DEFAULT_MAX_MODEL_STEPS,
140 token_estimate: crate::middleware::TokenEstimate::default(),
141 role: AgentRole::Main,
142 }
143 }
144
145 pub fn isolated_execution(mut self) -> Result<Self> {
150 self.sandbox = Arc::new(self.sandbox.isolated_execution()?);
151 Ok(self)
152 }
153
154 #[must_use]
156 pub fn session_id(mut self, session_id: impl Into<String>) -> Self {
157 self.session_id = session_id.into();
158 self
159 }
160
161 #[must_use]
163 pub fn session_context(mut self, context: SessionContext) -> Self {
164 self.session_context = context;
165 self
166 }
167
168 #[must_use]
170 pub fn middleware(mut self, middleware: MiddlewareStack) -> Self {
171 self.middleware = middleware;
172 self
173 }
174
175 #[must_use]
177 pub fn catalog_visible(mut self, visible: bool) -> Self {
178 self.catalog_visible = visible;
179 self
180 }
181
182 #[must_use]
184 pub fn token_estimate(mut self, estimate: crate::middleware::TokenEstimate) -> Self {
185 self.token_estimate = estimate;
186 self
187 }
188
189 #[must_use]
191 pub fn context_window(mut self, context_window: i64) -> Self {
192 self.context_window = context_window;
193 self.default_context_window = context_window;
194 self
195 }
196
197 #[must_use]
200 pub fn initial_replay_batches(mut self, max_batches: usize) -> Self {
201 self.initial_replay_batches = max_batches;
202 self
203 }
204
205 #[must_use]
207 pub fn max_model_steps(mut self, max_steps: usize) -> Self {
208 self.max_model_steps = max_steps;
209 self
210 }
211
212 #[must_use]
214 pub fn role(mut self, role: AgentRole) -> Self {
215 self.role = role;
216 self
217 }
218
219 #[must_use]
224 pub fn metadata(mut self, metadata: BTreeMap<String, Value>) -> Self {
225 self.metadata = metadata;
226 self.metadata_configured = true;
227 self
228 }
229
230 #[must_use]
236 pub fn usage_observer(
237 mut self,
238 observer: impl for<'a> Fn(&'a str, &'a TokenUsage) -> crate::BoxFuture<'a, Result<()>>
239 + Send
240 + Sync
241 + 'static,
242 ) -> Self {
243 self.usage_observer = Some(Arc::new(observer));
244 self
245 }
246
247 pub fn model_route(mut self, route: &str, reasoning_effort: Option<&str>) -> Result<Self> {
252 self.select_model_with_reasoning(route, reasoning_effort)?;
253 self.model_route_configured = true;
254 Ok(self)
255 }
256
257 #[must_use]
259 pub fn override_saved_model_route(mut self) -> Self {
260 self.model_route_configured = true;
261 self
262 }
263
264 fn select_model(&mut self, route: &str) -> Result<ModelChangedEvent> {
265 let default_window = self.default_context_window;
266 let choice = self.select_model_with_reasoning(route, None)?;
267 Ok(ModelChangedEvent {
268 route: choice.route.clone(),
269 model: choice.model.clone(),
270 reasoning_effort: choice.reasoning_effort.clone(),
271 model_context_window: Some(choice.context_window.unwrap_or(default_window)),
272 })
273 }
274
275 fn select_model_with_reasoning(
276 &mut self,
277 route: &str,
278 reasoning_effort: Option<&str>,
279 ) -> Result<&ModelChoice> {
280 let choice = self.model.resolve_choice(route, reasoning_effort)?;
281 self.provider.clone_from(&choice.route);
282 self.context_window = choice.context_window.unwrap_or(self.default_context_window);
283 Ok(choice)
284 }
285}
286
287#[derive(Clone)]
297pub struct AgentSender {
298 ingress: Arc<SubmissionIngress>,
299}
300
301#[derive(Debug, Clone, Copy, PartialEq, Eq)]
303pub enum MessageAcceptance {
304 Accepted,
306 AlreadyAccepted,
308}
309
310#[derive(Debug)]
312#[must_use = "wait for durable admission, or use AgentSender::send without an acknowledgment"]
313pub struct MessageAdmission {
314 receiver: oneshot::Receiver<Result<MessageAcceptance>>,
315}
316
317impl MessageAdmission {
318 pub async fn wait(self) -> Result<MessageAcceptance> {
323 self.receiver
324 .await
325 .map_err(|_| Error::Stopped("agent stopped before message admission".into()))?
326 }
327}
328
329type AdmissionReply = oneshot::Sender<Result<MessageAcceptance>>;
330
331#[derive(Clone)]
333pub struct WeakAgentSender {
334 ingress: Weak<SubmissionIngress>,
335}
336
337#[derive(Debug)]
339pub struct ValidatedSubmission(Submission);
340
341impl ValidatedSubmission {
342 pub fn new(submission: Submission) -> Result<Self> {
346 validate_submission(&submission)?;
347 Ok(Self(submission))
348 }
349
350 #[must_use]
352 pub fn submission(&self) -> &Submission {
353 &self.0
354 }
355
356 #[must_use]
358 pub fn into_submission(self) -> Submission {
359 self.0
360 }
361}
362
363impl AgentSender {
364 #[must_use]
366 pub fn downgrade(&self) -> WeakAgentSender {
367 WeakAgentSender {
368 ingress: Arc::downgrade(&self.ingress),
369 }
370 }
371
372 pub fn send(&self, submission: Submission) -> Result<()> {
377 self.send_validated(ValidatedSubmission::new(submission)?)
378 }
379
380 pub fn send_validated(&self, submission: ValidatedSubmission) -> Result<()> {
384 self.enqueue(submission, None)
385 }
386
387 pub fn send_with_admission(&self, submission: Submission) -> Result<MessageAdmission> {
394 self.send_validated_with_admission(ValidatedSubmission::new(submission)?)
395 }
396
397 pub fn send_validated_with_admission(
401 &self,
402 submission: ValidatedSubmission,
403 ) -> Result<MessageAdmission> {
404 if !matches!(submission.submission().op, Op::Message { .. }) {
405 return Err(Error::Config(
406 "admission acknowledgments require a message".into(),
407 ));
408 }
409 let (reply, receiver) = oneshot::channel();
410 self.enqueue(submission, Some(reply))?;
411 Ok(MessageAdmission { receiver })
412 }
413
414 fn enqueue(
415 &self,
416 submission: ValidatedSubmission,
417 admission: Option<AdmissionReply>,
418 ) -> Result<()> {
419 let submission = submission.0;
420 let mut last_sequence = self
421 .ingress
422 .last_sequence
423 .lock()
424 .map_err(|_| Error::Stopped("agent submission queue poisoned".into()))?;
425 let sequence = last_sequence
426 .checked_add(1)
427 .ok_or_else(|| Error::Busy("agent submission sequence exhausted".into()))?;
428 self.ingress
429 .sender
430 .try_send(SequencedSubmission {
431 sequence,
432 submission,
433 admission,
434 })
435 .map_err(|error| match error {
436 mpsc::error::TrySendError::Full(_) => {
437 Error::Busy("agent submission queue is full".into())
438 }
439 mpsc::error::TrySendError::Closed(_) => {
440 Error::Stopped("agent submission channel closed".into())
441 }
442 })?;
443 *last_sequence = sequence;
444 Ok(())
445 }
446
447 pub fn submit(&self, op: Op) -> Result<String> {
452 let id = Uuid::new_v4().to_string();
453 self.send(Submission { id: id.clone(), op })?;
454 Ok(id)
455 }
456}
457
458impl WeakAgentSender {
459 #[must_use]
461 pub fn upgrade(&self) -> Option<AgentSender> {
462 self.ingress
463 .upgrade()
464 .map(|ingress| AgentSender { ingress })
465 }
466}
467
468struct SubmissionIngress {
469 sender: mpsc::Sender<SequencedSubmission>,
470 last_sequence: Arc<Mutex<u64>>,
471}
472
473struct SequencedSubmission {
474 sequence: u64,
475 submission: Submission,
476 admission: Option<AdmissionReply>,
477}
478
479struct ReceivedSubmission {
480 submission: Submission,
481 admission: Option<AdmissionReply>,
482}
483
484struct SubmissionInbox {
485 receiver: mpsc::Receiver<SequencedSubmission>,
486 last_sent_sequence: Arc<Mutex<u64>>,
487 last_sequence: u64,
488}
489
490impl SubmissionInbox {
491 async fn recv(&mut self) -> Option<ReceivedSubmission> {
492 let queued = self.receiver.recv().await?;
493 self.last_sequence = queued.sequence;
494 Some(ReceivedSubmission {
495 submission: queued.submission,
496 admission: queued.admission,
497 })
498 }
499
500 fn cutoff(&self) -> Result<u64> {
501 self.last_sent_sequence
502 .lock()
503 .map(|sequence| *sequence)
504 .map_err(|_| Error::Stopped("agent submission queue poisoned".into()))
505 }
506}
507
508fn submission_channel(capacity: usize) -> (AgentSender, SubmissionInbox) {
509 let (sender, receiver) = mpsc::channel(capacity);
510 let last_sequence = Arc::new(Mutex::new(0));
511 let ingress = Arc::new(SubmissionIngress {
512 sender,
513 last_sequence: Arc::clone(&last_sequence),
514 });
515 (
516 AgentSender { ingress },
517 SubmissionInbox {
518 receiver,
519 last_sent_sequence: last_sequence,
520 last_sequence: 0,
521 },
522 )
523}
524
525#[cfg(test)]
526pub(crate) fn test_sender() -> WeakAgentSender {
527 submission_channel(1).0.downgrade()
528}
529
530pub fn validate_submission(submission: &Submission) -> Result<()> {
535 crate::validate_identifier("submission ID", &submission.id, crate::MAX_IDENTIFIER_BYTES)?;
536 match &submission.op {
537 Op::Message { message } => message.validate(session_file_limits()),
538 Op::Interrupt { turn_id } => {
539 crate::validate_identifier("turn ID", turn_id, crate::MAX_IDENTIFIER_BYTES)
540 }
541 Op::ExecApproval { id, .. } => {
542 crate::validate_identifier("approval ID", id, crate::MAX_IDENTIFIER_BYTES)
543 }
544 Op::CapabilityCommand {
545 capability,
546 command,
547 arguments,
548 input,
549 target,
550 } => {
551 crate::validate_identifier("capability ID", capability, MAX_OPERATION_BYTES)?;
552 crate::validate_identifier("command", command, MAX_OPERATION_BYTES)?;
553 if arguments.len() > crate::protocol::MAX_CAPABILITY_INPUT_BYTES {
554 return Err(Error::Config(
555 "middleware command arguments exceed size limit".into(),
556 ));
557 }
558 if input
559 .as_ref()
560 .is_some_and(|input| input.len() > crate::protocol::MAX_CAPABILITY_INPUT_BYTES)
561 {
562 return Err(Error::Config(
563 "middleware command input exceeds size limit".into(),
564 ));
565 }
566 if target.is_some_and(|target| target.batch_item_count == 0) {
567 return Err(Error::Config(
568 "message target item count must be positive".into(),
569 ));
570 }
571 Ok(())
572 }
573 Op::SetModel { route } => {
574 crate::validate_identifier("model route", route, crate::MAX_IDENTIFIER_BYTES)
575 }
576 Op::ResumeSession { session_id } => {
577 crate::validate_identifier("session ID", session_id, crate::MAX_IDENTIFIER_BYTES)
578 }
579 }
580}
581
582pub struct Agent {
584 sender: AgentSender,
585 events: mpsc::Receiver<JournalEvent>,
586 model_router: Arc<ModelRouter>,
587 frontend: FrontendExtensions,
588 frontend_sink: crate::middleware::FrontendEventSink,
589 session: SessionConfiguredEvent,
590 model: ModelInfo,
591 tool_count: usize,
592 next_before_sequence: Option<u64>,
593}
594
595impl Agent {
596 #[must_use]
598 pub fn sender(&self) -> AgentSender {
599 self.sender.clone()
600 }
601
602 #[must_use]
604 pub fn model_router(&self) -> Arc<ModelRouter> {
605 Arc::clone(&self.model_router)
606 }
607
608 pub async fn next_event(&mut self) -> Option<Event> {
612 self.events.recv().await.map(|record| record.event)
613 }
614
615 #[must_use]
617 pub fn frontend(&self) -> &FrontendExtensions {
618 &self.frontend
619 }
620
621 #[must_use]
624 pub fn frontend_sink(&self) -> crate::middleware::FrontendEventSink {
625 Arc::clone(&self.frontend_sink)
626 }
627
628 #[must_use]
630 pub fn session(&self) -> &SessionConfiguredEvent {
631 &self.session
632 }
633
634 #[must_use]
636 pub fn model(&self) -> &ModelInfo {
637 &self.model
638 }
639
640 #[must_use]
642 pub fn model_route(&self) -> &str {
643 &self.session.model.route
644 }
645
646 pub fn model_choices(
648 &self,
649 ) -> impl DoubleEndedIterator<Item = &ModelChoice> + ExactSizeIterator {
650 self.model_router.choices()
651 }
652
653 #[must_use]
655 pub const fn tool_count(&self) -> usize {
656 self.tool_count
657 }
658
659 #[must_use]
661 pub const fn next_before_sequence(&self) -> Option<u64> {
662 self.next_before_sequence
663 }
664
665 #[must_use]
669 pub fn into_parts(self) -> (AgentSender, AgentEvents) {
670 (self.sender, AgentEvents { inner: self.events })
671 }
672
673 #[must_use]
675 pub fn into_recorded_parts(
676 self,
677 ) -> (
678 AgentSender,
679 mpsc::Receiver<JournalEvent>,
680 SessionConfiguredEvent,
681 FrontendExtensions,
682 ) {
683 (self.sender, self.events, self.session, self.frontend)
684 }
685}
686
687pub struct AgentEvents {
689 inner: mpsc::Receiver<JournalEvent>,
690}
691
692impl AgentEvents {
693 pub async fn recv(&mut self) -> Option<Event> {
695 self.inner.recv().await.map(|record| record.event)
696 }
697
698 pub fn try_recv(&mut self) -> std::result::Result<Event, mpsc::error::TryRecvError> {
703 self.inner.try_recv().map(|record| record.event)
704 }
705}
706
707struct Runner {
708 config: AgentConfig,
709 runtime: Arc<RuntimeContext>,
710 system_prompt: Arc<str>,
711 catalog: Arc<Catalog>,
712 state: LiveCheckpoint,
713 transcript_delta: Vec<Arc<Value>>,
714 pending_session_start_stop: Option<String>,
715 turn_end_turn_id: Option<String>,
716 events: EventRecorder,
717}
718
719struct LiveCheckpoint(Arc<Checkpoint>);
721
722impl std::ops::Deref for LiveCheckpoint {
723 type Target = Checkpoint;
724
725 fn deref(&self) -> &Checkpoint {
726 &self.0
727 }
728}
729
730impl LiveCheckpoint {
731 fn make_mut(&mut self) -> &mut Checkpoint {
732 Arc::make_mut(&mut self.0)
733 }
734}
735
736impl Runner {
737 async fn run(&mut self, mut inbox: SubmissionInbox) -> Result<()> {
738 self.stop_resumed_turn_at_session_start().await?;
739 if let Some(pending) = self.state.pending_approval.clone() {
740 let submission_id = pending.submission_id.clone();
741 if let Err(error) = self.resume_pending(&mut inbox, pending).await {
742 self.fail_turn(&submission_id, error).await?;
743 }
744 } else if let Some(active) = self.state.active_execution.as_ref() {
745 let submission_id = active.submission_id.clone();
746 let turn_id = active.turn_id.clone();
747 if let Err(error) = self
748 .continue_turn(&mut inbox, submission_id.clone(), turn_id)
749 .await
750 {
751 self.fail_turn(&submission_id, error).await?;
752 }
753 }
754 loop {
755 if let Some(message) = self
756 .config
757 .middleware
758 .next_turn(&mut self.state.make_mut().pending_messages)?
759 {
760 let submission_id = message.submission_id.clone();
761 if let Err(error) = self.start_message_turn(&mut inbox, message).await {
762 self.fail_turn(&submission_id, error).await?;
763 }
764 continue;
765 }
766 let submission = {
767 let Some(submission) = inbox.recv().await else {
768 return Ok(());
769 };
770 submission
771 };
772 let ReceivedSubmission {
773 submission,
774 admission,
775 } = submission;
776 match submission.op {
777 Op::Message { message } => {
778 self.admit_idle_message(submission.id, message, admission)
779 .await?;
780 }
781 Op::Interrupt { .. } => {
782 self.emit(
783 submission.id,
784 EventMsg::Warning(WarningEvent {
785 message: "no active turn to interrupt".into(),
786 }),
787 )
788 .await?;
789 }
790 Op::ExecApproval { .. } => {
791 self.emit(
792 submission.id,
793 EventMsg::Warning(WarningEvent {
794 message: "no approval request is active".into(),
795 }),
796 )
797 .await?;
798 }
799 Op::CapabilityCommand {
800 capability,
801 command,
802 arguments,
803 input,
804 target,
805 } => {
806 self.capability_command(
807 submission.id,
808 capability,
809 command,
810 arguments,
811 input,
812 target,
813 )
814 .await?;
815 }
816 Op::SetModel { route } => {
817 self.set_model(submission.id, route).await?;
818 }
819 Op::ResumeSession { session_id } => {
820 self.request_resume(submission.id, session_id).await?;
821 }
822 }
823 }
824 }
825
826 async fn set_model(&mut self, submission_id: String, route: String) -> Result<()> {
827 let choice = match self.config.select_model(&route) {
828 Ok(choice) => choice,
829 Err(error) => {
830 self.emit(
831 submission_id,
832 EventMsg::Warning(WarningEvent {
833 message: error.to_string(),
834 }),
835 )
836 .await?;
837 return Ok(());
838 }
839 };
840 let active_route = choice.route.clone();
841 let active_model = choice.model.clone();
842 self.state.make_mut().model_route = Some(choice.route.clone());
843 self.persist_with_events(
844 vec![Event {
845 submission_id: Some(submission_id.into()),
846 msg: EventMsg::ModelChanged(choice),
847 }],
848 None,
849 )
850 .await?;
851 let runtime = Arc::make_mut(&mut self.runtime);
852 runtime.model_route = active_route;
853 runtime.model = active_model;
854 Ok(())
855 }
856
857 async fn request_resume(&self, submission_id: String, session_id: String) -> Result<()> {
858 let result = async {
859 if session_id.trim().is_empty() {
860 return Err(Error::Config("session ID cannot be empty".into()));
861 }
862 let checkpoint = self
863 .config
864 .checkpoints
865 .load(&session_id)
866 .await?
867 .ok_or_else(|| Error::Unknown(format!("session `{session_id}`")))?;
868 if checkpoint.version != CHECKPOINT_VERSION || checkpoint.session_id != session_id {
869 return Err(Error::Checkpoint(
870 "checkpoint does not match the requested session".into(),
871 ));
872 }
873 Ok(checkpoint.session_context)
874 }
875 .await;
876 match result {
877 Ok(context) => self.emit(
878 submission_id,
879 EventMsg::SessionResumeRequested(SessionResumeRequestedEvent {
880 session_id,
881 context,
882 }),
883 ),
884 Err(error) => self.emit(
885 submission_id,
886 EventMsg::Warning(WarningEvent {
887 message: error.to_string(),
888 }),
889 ),
890 }
891 .await
892 }
893
894 async fn capability_command(
895 &mut self,
896 submission_id: String,
897 capability: String,
898 command: String,
899 arguments: String,
900 input: Option<String>,
901 target: Option<crate::protocol::MessageTarget>,
902 ) -> Result<()> {
903 let output = self
904 .config
905 .middleware
906 .command(
907 &capability,
908 MiddlewareCommandContext {
909 command: &command,
910 arguments: &arguments,
911 input: input.as_deref(),
912 target,
913 session_id: &self.config.session_id,
914 session_context: &self.config.session_context,
915 checkpoint: &self.state,
916 checkpoints: Arc::clone(&self.config.checkpoints),
917 },
918 )
919 .await
920 .map(|output| output.events);
921 match output {
922 Ok(events) => {
923 for event in events {
924 self.emit(submission_id.as_str(), EventMsg::Frontend(event))
925 .await?;
926 }
927 }
928 Err(error) => {
929 self.emit(
930 submission_id,
931 EventMsg::Warning(WarningEvent {
932 message: error.to_string(),
933 }),
934 )
935 .await?
936 }
937 }
938 Ok(())
939 }
940
941 async fn save(&mut self) -> Result<u64> {
942 self.persist(None).await
943 }
944
945 async fn persist(&mut self, execution: Option<&ExecutionRecord>) -> Result<u64> {
946 self.persist_with_events(Vec::new(), execution).await
947 }
948
949 pub(super) async fn persist_with_events(
950 &mut self,
951 events: Vec<Event>,
952 execution: Option<&ExecutionRecord>,
953 ) -> Result<u64> {
954 let previous_sequence = self.state.sequence;
955 self.state.make_mut().sequence = self
956 .state
957 .sequence
958 .checked_add(1)
959 .ok_or_else(|| Error::Checkpoint("checkpoint sequence overflow".into()))?;
960 if let Err(error) = self
961 .events
962 .save(
963 Arc::clone(&self.state.0),
964 &self.transcript_delta,
965 execution,
966 events,
967 )
968 .await
969 {
970 self.state.make_mut().sequence = previous_sequence;
971 return Err(error);
972 }
973 self.transcript_delta.clear();
974 Ok(self.state.sequence)
975 }
976
977 fn record_model_call(&mut self) -> Result<()> {
978 let active = self
979 .state
980 .make_mut()
981 .active_execution
982 .as_mut()
983 .ok_or_else(|| Error::Checkpoint("model called without an active execution".into()))?;
984 active.model_calls = active
985 .model_calls
986 .checked_add(1)
987 .ok_or_else(|| Error::Checkpoint("execution model-call count overflow".into()))?;
988 Ok(())
989 }
990
991 async fn record_usage(&mut self, route: &str, usage: &TokenUsage) -> Result<()> {
992 let state = self.state.make_mut();
993 let mut total_usage = state.total_usage.clone();
995 total_usage.checked_add(usage).ok_or_else(|| {
996 Error::Provider("provider token usage exceeds the supported range".into())
997 })?;
998 let active = state.active_execution.as_mut().ok_or_else(|| {
999 Error::Checkpoint("usage recorded without an active execution".into())
1000 })?;
1001 let mut execution_usage = active.usage.clone();
1003 execution_usage.checked_add(usage).ok_or_else(|| {
1004 Error::Provider("provider token usage exceeds the supported range".into())
1005 })?;
1006 if let Some(observer) = &self.config.usage_observer {
1007 observer(route, usage).await?;
1008 }
1009 state.total_usage = total_usage;
1010 active.usage = execution_usage;
1011 Ok(())
1012 }
1013
1014 fn record_tools(&mut self, tool_calls: u64, failed_tool_calls: u64) -> Result<()> {
1015 let active = self
1016 .state
1017 .make_mut()
1018 .active_execution
1019 .as_mut()
1020 .ok_or_else(|| {
1021 Error::Checkpoint("tools recorded without an active execution".into())
1022 })?;
1023 let tool_calls = active
1024 .tool_calls
1025 .checked_add(tool_calls)
1026 .ok_or_else(|| Error::Checkpoint("execution tool-call count overflow".into()))?;
1027 let failed_tool_calls = active
1028 .failed_tool_calls
1029 .checked_add(failed_tool_calls)
1030 .ok_or_else(|| Error::Checkpoint("execution failed-tool count overflow".into()))?;
1031 active.tool_calls = tool_calls;
1032 active.failed_tool_calls = failed_tool_calls;
1033 Ok(())
1034 }
1035
1036 fn finish_execution(&mut self, outcome: ExecutionOutcome) -> Result<ExecutionRecord> {
1037 self.state
1038 .make_mut()
1039 .finish_execution(outcome, unix_timestamp_ms()?)
1040 }
1041
1042 async fn finish_and_persist_execution(
1043 &mut self,
1044 outcome: ExecutionOutcome,
1045 events: Vec<Event>,
1046 ) -> Result<u64> {
1047 let active_execution = self.state.active_execution.clone();
1049 let execution_stats = self.state.execution_stats.clone();
1050 let execution = self.finish_execution(outcome)?;
1051 match self.persist_with_events(events, Some(&execution)).await {
1052 Ok(sequence) => Ok(sequence),
1053 Err(error) => {
1054 self.state.make_mut().active_execution = active_execution;
1055 self.state.make_mut().execution_stats = execution_stats;
1056 Err(error)
1057 }
1058 }
1059 }
1060
1061 fn push_context(&mut self, item: Value) {
1062 let item = Arc::new(item);
1063 Arc::make_mut(&mut self.state.make_mut().context).push(Arc::clone(&item));
1064 self.transcript_delta.push(item);
1065 }
1066
1067 fn extend_context(&mut self, items: Vec<Value>) {
1068 for item in items {
1069 if crate::middleware::delivery_once::accept(
1070 &mut self.state.make_mut().delivered_once,
1071 &item,
1072 ) {
1073 self.push_context(item);
1074 }
1075 }
1076 }
1077
1078 async fn emit(&self, submission_id: impl Into<Arc<str>>, msg: EventMsg) -> Result<()> {
1079 send_event(
1080 &self.events,
1081 Event {
1082 submission_id: Some(submission_id.into()),
1083 msg,
1084 },
1085 )
1086 .await
1087 }
1088}
1089
1090fn unix_timestamp_ms() -> Result<i64> {
1091 let elapsed = SystemTime::now()
1092 .duration_since(UNIX_EPOCH)
1093 .map_err(|_| Error::Checkpoint("system clock predates the Unix epoch".into()))?;
1094 i64::try_from(elapsed.as_millis())
1095 .map_err(|_| Error::Checkpoint("system clock exceeds the supported range".into()))
1096}
1097
1098async fn send_event(events: &EventRecorder, event: Event) -> Result<()> {
1099 events.record(event).await
1100}
1101
1102fn try_send_event(events: &EventRecorder, event: Event) -> Result<()> {
1103 events.try_record(event)
1104}
1105
1106#[cfg(test)]
1107#[path = "runtime_tests.rs"]
1108mod tests;