1use std::sync::Arc;
4
5use async_trait::async_trait;
6
7use crate::{ReviewRequest, Session, SessionError, SessionId};
8
9#[derive(Clone, Debug, Default, Eq, PartialEq)]
11pub enum CreateSessionMode {
12 #[default]
14 Regular,
15 Draft,
17 Orchestrator,
19 OrchestrationChild {
21 task_id: i64,
23 },
24 OrchestrationResearch {
27 task_id: i64,
29 },
30 Stacked {
32 parent_session_id: SessionId,
34 },
35}
36
37#[derive(Clone, Debug, Eq, PartialEq)]
39pub struct CreateSessionRequest {
40 pub inherit_from_session_id: Option<SessionId>,
44 pub mode: CreateSessionMode,
46 pub project_id: i64,
48}
49
50#[derive(Clone, Debug, Eq, PartialEq)]
52pub struct QuestionAnswer {
53 pub answer: String,
55 pub question: String,
57}
58
59#[derive(Clone, Debug, Eq, PartialEq)]
61pub struct AnswerQuestionsRequest {
62 pub answers: Vec<QuestionAnswer>,
64}
65
66#[derive(Clone, Debug, Eq, PartialEq)]
68pub struct CoordinatorMessageRequest {
69 pub message: String,
71 pub operation_id: String,
73 pub visibility: CoordinatorMessageVisibility,
75}
76
77#[derive(Clone, Copy, Debug, Default, Eq, PartialEq)]
79pub enum CoordinatorMessageVisibility {
80 #[default]
82 Hidden,
83 Visible,
85}
86
87#[async_trait]
92pub trait SessionBackend: Send + Sync {
93 async fn create_session(
95 &self,
96 request: CreateSessionRequest,
97 ) -> Result<SessionId, SessionError>;
98
99 async fn get_session(&self, session_id: &SessionId) -> Result<Option<Session>, SessionError>;
101
102 async fn send_message(
104 &self,
105 session_id: &SessionId,
106 message: String,
107 ) -> Result<(), SessionError>;
108
109 async fn submit_coordinator_message(
112 &self,
113 session_id: &SessionId,
114 request: CoordinatorMessageRequest,
115 ) -> Result<(), SessionError>;
116
117 async fn answer_questions(
119 &self,
120 session_id: &SessionId,
121 request: AnswerQuestionsRequest,
122 ) -> Result<(), SessionError>;
123
124 async fn cancel_session(&self, session_id: &SessionId) -> Result<(), SessionError>;
126
127 async fn merge_session(&self, session_id: &SessionId) -> Result<(), SessionError>;
129
130 async fn create_review_request(
133 &self,
134 session_id: &SessionId,
135 ) -> Result<ReviewRequest, SessionError>;
136}
137
138#[derive(Clone)]
140pub struct SessionService {
141 backend: Arc<dyn SessionBackend>,
142}
143
144impl SessionService {
145 pub fn new(backend: Arc<dyn SessionBackend>) -> Self {
147 Self { backend }
148 }
149
150 pub async fn create_session(
155 &self,
156 request: CreateSessionRequest,
157 ) -> Result<SessionId, SessionError> {
158 self.backend.create_session(request).await
159 }
160
161 pub async fn get_session(
166 &self,
167 session_id: &SessionId,
168 ) -> Result<Option<Session>, SessionError> {
169 self.backend.get_session(session_id).await
170 }
171
172 pub async fn send_message(
177 &self,
178 session_id: &SessionId,
179 message: impl Into<String> + Send,
180 ) -> Result<(), SessionError> {
181 self.backend.send_message(session_id, message.into()).await
182 }
183
184 pub async fn submit_coordinator_message(
189 &self,
190 session_id: &SessionId,
191 request: CoordinatorMessageRequest,
192 ) -> Result<(), SessionError> {
193 self.backend
194 .submit_coordinator_message(session_id, request)
195 .await
196 }
197
198 pub async fn answer_questions(
204 &self,
205 session_id: &SessionId,
206 request: AnswerQuestionsRequest,
207 ) -> Result<(), SessionError> {
208 self.backend.answer_questions(session_id, request).await
209 }
210
211 pub async fn cancel_session(&self, session_id: &SessionId) -> Result<(), SessionError> {
217 self.backend.cancel_session(session_id).await
218 }
219
220 pub async fn merge_session(&self, session_id: &SessionId) -> Result<(), SessionError> {
225 self.backend.merge_session(session_id).await
226 }
227
228 pub async fn create_review_request(
235 &self,
236 session_id: &SessionId,
237 ) -> Result<ReviewRequest, SessionError> {
238 self.backend.create_review_request(session_id).await
239 }
240}
241
242#[cfg(test)]
243mod tests {
244 use std::collections::VecDeque;
245 use std::sync::Mutex;
246
247 use ag_agent::{AgentKind, AgentModel, AgentSelection, ReasoningLevel, SpeedMode};
248 use ag_forge::{ForgeKind, ReviewRequestState, ReviewRequestSummary};
249
250 use super::*;
251 use crate::{SessionMessage, SessionMessageKind, SessionRole, SessionSettings, SessionStatus};
252
253 #[derive(Default)]
254 struct FakeBackend {
255 state: Mutex<FakeBackendState>,
256 }
257
258 impl FakeBackend {
259 fn from_state(state: FakeBackendState) -> Self {
260 Self {
261 state: Mutex::new(state),
262 }
263 }
264
265 fn calls(&self) -> Vec<String> {
266 self.state
267 .lock()
268 .map(|state| state.calls.clone())
269 .unwrap_or_default()
270 }
271 }
272
273 #[derive(Default)]
274 struct FakeBackendState {
275 calls: Vec<String>,
276 create_results: VecDeque<Result<SessionId, SessionError>>,
277 get_result: Option<Result<Option<Session>, SessionError>>,
278 review_result: Option<Result<ReviewRequest, SessionError>>,
279 unit_results: VecDeque<Result<(), SessionError>>,
280 }
281
282 #[async_trait]
283 impl SessionBackend for FakeBackend {
284 async fn create_session(
285 &self,
286 request: CreateSessionRequest,
287 ) -> Result<SessionId, SessionError> {
288 let mut state = self
289 .state
290 .lock()
291 .expect("fake backend state should remain available");
292 state.calls.push(format!("create:{:?}", request.mode));
293
294 state
295 .create_results
296 .pop_front()
297 .unwrap_or_else(|| Err(SessionError::Operation("missing result".to_string())))
298 }
299
300 async fn get_session(
301 &self,
302 _session_id: &SessionId,
303 ) -> Result<Option<Session>, SessionError> {
304 self.state
305 .lock()
306 .expect("fake backend state should remain available")
307 .get_result
308 .clone()
309 .unwrap_or_else(|| Err(SessionError::Operation("missing result".to_string())))
310 }
311
312 async fn send_message(
313 &self,
314 session_id: &SessionId,
315 message: String,
316 ) -> Result<(), SessionError> {
317 let mut state = self
318 .state
319 .lock()
320 .expect("fake backend state should remain available");
321 state.calls.push(format!("send:{session_id}:{message}"));
322
323 state
324 .unit_results
325 .pop_front()
326 .unwrap_or_else(|| Err(SessionError::Operation("missing result".to_string())))
327 }
328
329 async fn submit_coordinator_message(
330 &self,
331 session_id: &SessionId,
332 request: CoordinatorMessageRequest,
333 ) -> Result<(), SessionError> {
334 let mut state = self
335 .state
336 .lock()
337 .expect("fake backend state should remain available");
338 state.calls.push(format!(
339 "submit-coordinator:{session_id}:{}:{}",
340 request.operation_id, request.message
341 ));
342
343 state
344 .unit_results
345 .pop_front()
346 .unwrap_or_else(|| Err(SessionError::Operation("missing result".to_string())))
347 }
348
349 async fn answer_questions(
350 &self,
351 session_id: &SessionId,
352 request: AnswerQuestionsRequest,
353 ) -> Result<(), SessionError> {
354 let mut state = self
355 .state
356 .lock()
357 .expect("fake backend state should remain available");
358 state
359 .calls
360 .push(format!("answer:{session_id}:{}", request.answers.len()));
361
362 state
363 .unit_results
364 .pop_front()
365 .unwrap_or_else(|| Err(SessionError::Operation("missing result".to_string())))
366 }
367
368 async fn cancel_session(&self, session_id: &SessionId) -> Result<(), SessionError> {
369 let mut state = self
370 .state
371 .lock()
372 .expect("fake backend state should remain available");
373 state.calls.push(format!("cancel:{session_id}"));
374
375 state
376 .unit_results
377 .pop_front()
378 .unwrap_or_else(|| Err(SessionError::Operation("missing result".to_string())))
379 }
380
381 async fn merge_session(&self, session_id: &SessionId) -> Result<(), SessionError> {
382 let mut state = self
383 .state
384 .lock()
385 .expect("fake backend state should remain available");
386 state.calls.push(format!("merge:{session_id}"));
387
388 state
389 .unit_results
390 .pop_front()
391 .unwrap_or_else(|| Err(SessionError::Operation("missing result".to_string())))
392 }
393
394 async fn create_review_request(
395 &self,
396 session_id: &SessionId,
397 ) -> Result<ReviewRequest, SessionError> {
398 let mut state = self
399 .state
400 .lock()
401 .expect("fake backend state should remain available");
402 state.calls.push(format!("review:{session_id}"));
403
404 state
405 .review_result
406 .clone()
407 .unwrap_or_else(|| Err(SessionError::Operation("missing result".to_string())))
408 }
409 }
410
411 fn session_fixture() -> Session {
412 Session {
413 created_at: 10,
414 draft_prompt: None,
415 id: SessionId::from("session-1"),
416 messages: vec![SessionMessage::new(
417 0,
418 SessionMessageKind::UserPrompt,
419 "build it",
420 )],
421 published_upstream_ref: None,
422 questions: Vec::new(),
423 queued_messages: Vec::new(),
424 review_request: None,
425 settings: SessionSettings {
426 agent: AgentSelection::new(AgentKind::Codex, AgentModel::Gpt56Sol),
427 base_branch: "main".to_string(),
428 is_draft: false,
429 parent_session_id: None,
430 personality_id: Some("reviewer".to_string()),
431 project_id: 7,
432 reasoning_level: ReasoningLevel::High,
433 role: SessionRole::Worker,
434 speed_mode: SpeedMode::Normal,
435 },
436 status: SessionStatus::Review,
437 summary: Some("Implemented it".to_string()),
438 title: Some("Build it".to_string()),
439 updated_at: 20,
440 }
441 }
442
443 fn review_request_fixture() -> ReviewRequest {
444 ReviewRequest {
445 last_refreshed_at: 30,
446 summary: ReviewRequestSummary {
447 display_id: "#42".to_string(),
448 forge_kind: ForgeKind::GitHub,
449 source_branch: "wt/session-1".to_string(),
450 state: ReviewRequestState::Open,
451 status_summary: None,
452 target_branch: "main".to_string(),
453 title: "Build it".to_string(),
454 web_url: "https://example.test/pull/42".to_string(),
455 },
456 }
457 }
458
459 #[tokio::test]
460 async fn service_delegates_create_and_get() {
461 let expected_session = session_fixture();
463 let backend = Arc::new(FakeBackend::from_state(FakeBackendState {
464 create_results: VecDeque::from([Ok(SessionId::from("session-1"))]),
465 get_result: Some(Ok(Some(expected_session.clone()))),
466 ..FakeBackendState::default()
467 }));
468 let service = SessionService::new(backend.clone());
469
470 let session_id = service
472 .create_session(CreateSessionRequest {
473 inherit_from_session_id: None,
474 mode: CreateSessionMode::Regular,
475 project_id: 7,
476 })
477 .await
478 .expect("session should be created");
479 let loaded_session = service
480 .get_session(&session_id)
481 .await
482 .expect("session should load");
483
484 assert_eq!(loaded_session, Some(expected_session));
486 assert_eq!(backend.calls(), ["create:Regular"]);
487 }
488
489 #[tokio::test]
490 async fn service_delegates_mutating_operations_through_clones() {
491 let expected_review_request = review_request_fixture();
493 let backend = Arc::new(FakeBackend::from_state(FakeBackendState {
494 review_result: Some(Ok(expected_review_request.clone())),
495 unit_results: VecDeque::from([Ok(()), Ok(()), Ok(()), Ok(()), Ok(())]),
496 ..FakeBackendState::default()
497 }));
498 let session_id = SessionId::from("session-1");
499 let service = SessionService::new(backend.clone());
500 let cloned_service = service.clone();
501 let answers = AnswerQuestionsRequest {
502 answers: vec![QuestionAnswer {
503 answer: "main".to_string(),
504 question: "Which branch?".to_string(),
505 }],
506 };
507
508 service
510 .send_message(&session_id, "continue")
511 .await
512 .expect("message should be sent");
513 service
514 .submit_coordinator_message(
515 &session_id,
516 CoordinatorMessageRequest {
517 message: "roll up".to_string(),
518 operation_id: "rollup-7".to_string(),
519 visibility: CoordinatorMessageVisibility::Hidden,
520 },
521 )
522 .await
523 .expect("coordinator message should be submitted");
524 cloned_service
525 .answer_questions(&session_id, answers)
526 .await
527 .expect("questions should be answered");
528 service
529 .cancel_session(&session_id)
530 .await
531 .expect("cancel should be requested");
532 cloned_service
533 .merge_session(&session_id)
534 .await
535 .expect("merge should be requested");
536 let review_request = service
537 .create_review_request(&session_id)
538 .await
539 .expect("review request should be created");
540
541 assert_eq!(review_request, expected_review_request);
543 assert_eq!(
544 backend.calls(),
545 [
546 "send:session-1:continue",
547 "submit-coordinator:session-1:rollup-7:roll up",
548 "answer:session-1:1",
549 "cancel:session-1",
550 "merge:session-1",
551 "review:session-1"
552 ]
553 );
554 }
555
556 #[tokio::test]
557 async fn service_preserves_backend_errors() {
558 let expected_error = SessionError::Operation("cannot create".to_string());
560 let backend = Arc::new(FakeBackend::from_state(FakeBackendState {
561 create_results: VecDeque::from([Err(expected_error.clone())]),
562 ..FakeBackendState::default()
563 }));
564 let service = SessionService::new(backend);
565
566 let error = service
568 .create_session(CreateSessionRequest {
569 inherit_from_session_id: None,
570 mode: CreateSessionMode::Draft,
571 project_id: 7,
572 })
573 .await
574 .expect_err("backend error should be preserved");
575
576 assert_eq!(error, expected_error);
578 }
579
580 #[tokio::test]
581 async fn fake_backend_requires_explicit_results() {
582 let backend = Arc::new(FakeBackend::default());
584 let session_id = SessionId::from("session-1");
585 let service = SessionService::new(backend);
586
587 let create_error = service
589 .create_session(CreateSessionRequest {
590 inherit_from_session_id: None,
591 mode: CreateSessionMode::Regular,
592 project_id: 7,
593 })
594 .await
595 .expect_err("create should require a result");
596 let get_error = service
597 .get_session(&session_id)
598 .await
599 .expect_err("get should require a result");
600 let send_error = service
601 .send_message(&session_id, "continue")
602 .await
603 .expect_err("send should require a result");
604 let coordinator_error = service
605 .submit_coordinator_message(
606 &session_id,
607 CoordinatorMessageRequest {
608 message: "roll up".to_string(),
609 operation_id: "rollup-1".to_string(),
610 visibility: CoordinatorMessageVisibility::Hidden,
611 },
612 )
613 .await
614 .expect_err("coordinator submission should require a result");
615 let answer_error = service
616 .answer_questions(
617 &session_id,
618 AnswerQuestionsRequest {
619 answers: Vec::new(),
620 },
621 )
622 .await
623 .expect_err("answers should require a result");
624 let cancel_error = service
625 .cancel_session(&session_id)
626 .await
627 .expect_err("cancel should require a result");
628 let merge_error = service
629 .merge_session(&session_id)
630 .await
631 .expect_err("merge should require a result");
632 let review_error = service
633 .create_review_request(&session_id)
634 .await
635 .expect_err("review should require a result");
636 let errors = [
637 create_error,
638 get_error,
639 send_error,
640 coordinator_error,
641 answer_error,
642 cancel_error,
643 merge_error,
644 review_error,
645 ];
646
647 assert!(
649 errors
650 .into_iter()
651 .all(|error| error == SessionError::Operation("missing result".to_string()))
652 );
653 }
654}