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