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::{
252 PermissionMode, SessionMessage, SessionMessageKind, SessionRole, SessionSettings,
253 SessionStatus,
254 };
255
256 #[derive(Default)]
257 struct FakeBackend {
258 state: Mutex<FakeBackendState>,
259 }
260
261 impl FakeBackend {
262 fn from_state(state: FakeBackendState) -> Self {
263 Self {
264 state: Mutex::new(state),
265 }
266 }
267
268 fn calls(&self) -> Vec<String> {
269 self.state
270 .lock()
271 .map(|state| state.calls.clone())
272 .unwrap_or_default()
273 }
274 }
275
276 #[derive(Default)]
277 struct FakeBackendState {
278 calls: Vec<String>,
279 create_results: VecDeque<Result<SessionId, SessionError>>,
280 get_result: Option<Result<Option<Session>, SessionError>>,
281 review_result: Option<Result<ReviewRequest, SessionError>>,
282 unit_results: VecDeque<Result<(), SessionError>>,
283 }
284
285 #[async_trait]
286 impl SessionBackend for FakeBackend {
287 async fn create_session(
288 &self,
289 request: CreateSessionRequest,
290 ) -> Result<SessionId, SessionError> {
291 let mut state = self
292 .state
293 .lock()
294 .expect("fake backend state should remain available");
295 state.calls.push(format!("create:{:?}", request.mode));
296
297 state
298 .create_results
299 .pop_front()
300 .unwrap_or_else(|| Err(SessionError::Operation("missing result".to_string())))
301 }
302
303 async fn get_session(
304 &self,
305 _session_id: &SessionId,
306 ) -> Result<Option<Session>, SessionError> {
307 self.state
308 .lock()
309 .expect("fake backend state should remain available")
310 .get_result
311 .clone()
312 .unwrap_or_else(|| Err(SessionError::Operation("missing result".to_string())))
313 }
314
315 async fn send_message(
316 &self,
317 session_id: &SessionId,
318 message: String,
319 ) -> Result<(), SessionError> {
320 let mut state = self
321 .state
322 .lock()
323 .expect("fake backend state should remain available");
324 state.calls.push(format!("send:{session_id}:{message}"));
325
326 state
327 .unit_results
328 .pop_front()
329 .unwrap_or_else(|| Err(SessionError::Operation("missing result".to_string())))
330 }
331
332 async fn submit_coordinator_message(
333 &self,
334 session_id: &SessionId,
335 request: CoordinatorMessageRequest,
336 ) -> Result<(), SessionError> {
337 let mut state = self
338 .state
339 .lock()
340 .expect("fake backend state should remain available");
341 state.calls.push(format!(
342 "submit-coordinator:{session_id}:{}:{}",
343 request.operation_id, request.message
344 ));
345
346 state
347 .unit_results
348 .pop_front()
349 .unwrap_or_else(|| Err(SessionError::Operation("missing result".to_string())))
350 }
351
352 async fn answer_questions(
353 &self,
354 session_id: &SessionId,
355 request: AnswerQuestionsRequest,
356 ) -> Result<(), SessionError> {
357 let mut state = self
358 .state
359 .lock()
360 .expect("fake backend state should remain available");
361 state
362 .calls
363 .push(format!("answer:{session_id}:{}", request.answers.len()));
364
365 state
366 .unit_results
367 .pop_front()
368 .unwrap_or_else(|| Err(SessionError::Operation("missing result".to_string())))
369 }
370
371 async fn cancel_session(&self, session_id: &SessionId) -> Result<(), SessionError> {
372 let mut state = self
373 .state
374 .lock()
375 .expect("fake backend state should remain available");
376 state.calls.push(format!("cancel:{session_id}"));
377
378 state
379 .unit_results
380 .pop_front()
381 .unwrap_or_else(|| Err(SessionError::Operation("missing result".to_string())))
382 }
383
384 async fn merge_session(&self, session_id: &SessionId) -> Result<(), SessionError> {
385 let mut state = self
386 .state
387 .lock()
388 .expect("fake backend state should remain available");
389 state.calls.push(format!("merge:{session_id}"));
390
391 state
392 .unit_results
393 .pop_front()
394 .unwrap_or_else(|| Err(SessionError::Operation("missing result".to_string())))
395 }
396
397 async fn create_review_request(
398 &self,
399 session_id: &SessionId,
400 ) -> Result<ReviewRequest, SessionError> {
401 let mut state = self
402 .state
403 .lock()
404 .expect("fake backend state should remain available");
405 state.calls.push(format!("review:{session_id}"));
406
407 state
408 .review_result
409 .clone()
410 .unwrap_or_else(|| Err(SessionError::Operation("missing result".to_string())))
411 }
412 }
413
414 fn session_fixture() -> Session {
415 Session {
416 created_at: 10,
417 draft_prompt: None,
418 id: SessionId::from("session-1"),
419 messages: vec![SessionMessage::new(
420 0,
421 SessionMessageKind::UserPrompt,
422 "build it",
423 )],
424 published_upstream_ref: None,
425 questions: Vec::new(),
426 queued_messages: Vec::new(),
427 review_request: None,
428 settings: SessionSettings {
429 agent: AgentSelection::new(AgentKind::Codex, AgentModel::Gpt56Sol),
430 base_branch: "main".to_string(),
431 is_draft: false,
432 parent_session_id: None,
433 permission_mode: PermissionMode::AutoEdit,
434 personality_id: Some("reviewer".to_string()),
435 project_id: 7,
436 reasoning_level: ReasoningLevel::High,
437 role: SessionRole::Worker,
438 speed_mode: SpeedMode::Normal,
439 },
440 status: SessionStatus::Review,
441 summary: Some("Implemented it".to_string()),
442 title: Some("Build it".to_string()),
443 updated_at: 20,
444 }
445 }
446
447 fn review_request_fixture() -> ReviewRequest {
448 ReviewRequest {
449 last_refreshed_at: 30,
450 summary: ReviewRequestSummary {
451 display_id: "#42".to_string(),
452 forge_kind: ForgeKind::GitHub,
453 source_branch: "wt/session-1".to_string(),
454 state: ReviewRequestState::Open,
455 status_summary: None,
456 target_branch: "main".to_string(),
457 title: "Build it".to_string(),
458 web_url: "https://example.test/pull/42".to_string(),
459 },
460 }
461 }
462
463 #[tokio::test]
464 async fn service_delegates_create_and_get() {
465 let expected_session = session_fixture();
467 let backend = Arc::new(FakeBackend::from_state(FakeBackendState {
468 create_results: VecDeque::from([Ok(SessionId::from("session-1"))]),
469 get_result: Some(Ok(Some(expected_session.clone()))),
470 ..FakeBackendState::default()
471 }));
472 let service = SessionService::new(backend.clone());
473
474 let session_id = service
476 .create_session(CreateSessionRequest {
477 inherit_from_session_id: None,
478 mode: CreateSessionMode::Regular,
479 project_id: 7,
480 })
481 .await
482 .expect("session should be created");
483 let loaded_session = service
484 .get_session(&session_id)
485 .await
486 .expect("session should load");
487
488 assert_eq!(loaded_session, Some(expected_session));
490 assert_eq!(backend.calls(), ["create:Regular"]);
491 }
492
493 #[tokio::test]
494 async fn service_delegates_mutating_operations_through_clones() {
495 let expected_review_request = review_request_fixture();
497 let backend = Arc::new(FakeBackend::from_state(FakeBackendState {
498 review_result: Some(Ok(expected_review_request.clone())),
499 unit_results: VecDeque::from([Ok(()), Ok(()), Ok(()), Ok(()), Ok(())]),
500 ..FakeBackendState::default()
501 }));
502 let session_id = SessionId::from("session-1");
503 let service = SessionService::new(backend.clone());
504 let cloned_service = service.clone();
505 let answers = AnswerQuestionsRequest {
506 answers: vec![QuestionAnswer {
507 answer: "main".to_string(),
508 question: "Which branch?".to_string(),
509 }],
510 };
511
512 service
514 .send_message(&session_id, "continue")
515 .await
516 .expect("message should be sent");
517 service
518 .submit_coordinator_message(
519 &session_id,
520 CoordinatorMessageRequest {
521 message: "roll up".to_string(),
522 operation_id: "rollup-7".to_string(),
523 visibility: CoordinatorMessageVisibility::Hidden,
524 },
525 )
526 .await
527 .expect("coordinator message should be submitted");
528 cloned_service
529 .answer_questions(&session_id, answers)
530 .await
531 .expect("questions should be answered");
532 service
533 .cancel_session(&session_id)
534 .await
535 .expect("cancel should be requested");
536 cloned_service
537 .merge_session(&session_id)
538 .await
539 .expect("merge should be requested");
540 let review_request = service
541 .create_review_request(&session_id)
542 .await
543 .expect("review request should be created");
544
545 assert_eq!(review_request, expected_review_request);
547 assert_eq!(
548 backend.calls(),
549 [
550 "send:session-1:continue",
551 "submit-coordinator:session-1:rollup-7:roll up",
552 "answer:session-1:1",
553 "cancel:session-1",
554 "merge:session-1",
555 "review:session-1"
556 ]
557 );
558 }
559
560 #[tokio::test]
561 async fn service_preserves_backend_errors() {
562 let expected_error = SessionError::Operation("cannot create".to_string());
564 let backend = Arc::new(FakeBackend::from_state(FakeBackendState {
565 create_results: VecDeque::from([Err(expected_error.clone())]),
566 ..FakeBackendState::default()
567 }));
568 let service = SessionService::new(backend);
569
570 let error = service
572 .create_session(CreateSessionRequest {
573 inherit_from_session_id: None,
574 mode: CreateSessionMode::Draft,
575 project_id: 7,
576 })
577 .await
578 .expect_err("backend error should be preserved");
579
580 assert_eq!(error, expected_error);
582 }
583
584 #[tokio::test]
585 async fn fake_backend_requires_explicit_results() {
586 let backend = Arc::new(FakeBackend::default());
588 let session_id = SessionId::from("session-1");
589 let service = SessionService::new(backend);
590
591 let create_error = service
593 .create_session(CreateSessionRequest {
594 inherit_from_session_id: None,
595 mode: CreateSessionMode::Regular,
596 project_id: 7,
597 })
598 .await
599 .expect_err("create should require a result");
600 let get_error = service
601 .get_session(&session_id)
602 .await
603 .expect_err("get should require a result");
604 let send_error = service
605 .send_message(&session_id, "continue")
606 .await
607 .expect_err("send should require a result");
608 let coordinator_error = service
609 .submit_coordinator_message(
610 &session_id,
611 CoordinatorMessageRequest {
612 message: "roll up".to_string(),
613 operation_id: "rollup-1".to_string(),
614 visibility: CoordinatorMessageVisibility::Hidden,
615 },
616 )
617 .await
618 .expect_err("coordinator submission should require a result");
619 let answer_error = service
620 .answer_questions(
621 &session_id,
622 AnswerQuestionsRequest {
623 answers: Vec::new(),
624 },
625 )
626 .await
627 .expect_err("answers should require a result");
628 let cancel_error = service
629 .cancel_session(&session_id)
630 .await
631 .expect_err("cancel should require a result");
632 let merge_error = service
633 .merge_session(&session_id)
634 .await
635 .expect_err("merge should require a result");
636 let review_error = service
637 .create_review_request(&session_id)
638 .await
639 .expect_err("review should require a result");
640 let errors = [
641 create_error,
642 get_error,
643 send_error,
644 coordinator_error,
645 answer_error,
646 cancel_error,
647 merge_error,
648 review_error,
649 ];
650
651 assert!(
653 errors
654 .into_iter()
655 .all(|error| error == SessionError::Operation("missing result".to_string()))
656 );
657 }
658}