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