diff --git a/src/config.rs b/src/config.rs index 0898fcf..09e1850 100644 --- a/src/config.rs +++ b/src/config.rs @@ -762,7 +762,7 @@ rauthy: assert_eq!(cfg.url, "http://0.0.0.0:5500"); assert_eq!(cfg.hiqlite_raft_port, 9100); assert_eq!(cfg.hiqlite_api_port, 9200); - assert_eq!(cfg.rauthy_enabled, true); + assert!(cfg.rauthy_enabled); assert_eq!(cfg.rauthy_port, 4444); clear_ofm_env(); diff --git a/src/opencode_sdk/client.rs b/src/opencode_sdk/client.rs index ab3b78d..756a652 100644 --- a/src/opencode_sdk/client.rs +++ b/src/opencode_sdk/client.rs @@ -731,7 +731,7 @@ mod tests { let client = reqwest::Client::new(); let resp = client - .get(&format!("http://127.0.0.1:{port}/event")) + .get(format!("http://127.0.0.1:{port}/event")) .send() .await .unwrap(); diff --git a/src/providers/opencode_sdk_provider.rs b/src/providers/opencode_sdk_provider.rs index 842cafa..ab9a1f3 100644 --- a/src/providers/opencode_sdk_provider.rs +++ b/src/providers/opencode_sdk_provider.rs @@ -1,10 +1,12 @@ use std::path::{Path, PathBuf}; +use std::sync::Arc; use std::sync::Mutex; use std::time::Duration; use async_trait::async_trait; use futures_util::StreamExt; use tokio::sync::mpsc; +use tokio::sync::Notify; use uuid::Uuid; use crate::opencode_sdk::client::EventStreamCancellation; @@ -26,8 +28,10 @@ pub struct OpenCodeSdkProvider { /// Last known session id — used by `abort_turn` for the best-effort /// `client.session.abort` call. session_id: Mutex>, - /// Cancellation handle for the in-flight event stream reader task. + /// Cancellation handle for the in-flight event stream subscription. event_cancellation: Mutex>, + /// Notify channel to signal the reader task to exit. + reader_cancellation: Mutex>>, /// User id used to key the pool. May be `None` for one-shot operations /// (`get_models_list`, `one_shot_prompt`, title generation) which /// spawn transient servers outside the pool. @@ -55,6 +59,7 @@ impl OpenCodeSdkProvider { client: Mutex::new(None), session_id: Mutex::new(None), event_cancellation: Mutex::new(None), + reader_cancellation: Mutex::new(None), user_id: Mutex::new(None), working_dir: Mutex::new(None), log_data, @@ -68,6 +73,30 @@ impl OpenCodeSdkProvider { *self.user_id.lock().unwrap() = Some(user_id); } + fn client_and_user_id(&self) -> Result<(OpencodeClient, Uuid), ProviderError> { + let client = self + .client + .lock() + .unwrap() + .clone() + .ok_or(ProviderError::NotStarted)?; + let user_id = self + .user_id + .lock() + .unwrap() + .ok_or_else(|| ProviderError::Protocol("user_id not set on provider".into()))?; + Ok((client, user_id)) + } + + fn cancel_inflight(&self) { + if let Some(cancel) = self.reader_cancellation.lock().unwrap().take() { + cancel.notify_one(); + } + if let Some(cancellation) = self.event_cancellation.lock().unwrap().take() { + cancellation.cancel(); + } + } + fn build_server_config(&self) -> serde_json::Value { let mut base = serde_json::json!({ "provider": {}, @@ -133,6 +162,8 @@ impl OpenCodeSdkProvider { session_id: &str, user_id: Uuid, ) -> Result, ProviderError> { + self.cancel_inflight(); + tracing::info!( session_id = %session_id, "Subscribing to opencode global event stream" @@ -147,6 +178,9 @@ impl OpenCodeSdkProvider { let cancellation = event_stream.cancellation_handle(); *self.event_cancellation.lock().unwrap() = Some(cancellation); + let reader_stop = Arc::new(Notify::new()); + *self.reader_cancellation.lock().unwrap() = Some(reader_stop.clone()); + let s_id = session_id.to_string(); let (tx, rx) = mpsc::channel(1024); @@ -191,7 +225,7 @@ impl OpenCodeSdkProvider { error = %e, "Event stream error" ); - let _ = tx + let _ = tx .send(ProviderEvent::Error { error: e.to_string(), timestamp: chrono::Utc::now().naive_utc(), @@ -206,6 +240,13 @@ impl OpenCodeSdkProvider { pool.update_timestamp(user_id).await; } + _ = reader_stop.notified() => { + tracing::info!( + session_id = %s_id, + "Event reader task cancelled" + ); + break; + } } } tracing::info!(session_id = %s_id, "Event reader task exited"); @@ -234,6 +275,7 @@ fn map_sdk_event_to_provider_event( global: &GlobalEvent, session_id: &str, ) -> Option { + let now = chrono::Utc::now().naive_utc(); match &global.payload { Event::MessagePartUpdated(data) => match &data.part { Part::Text(t) => Some(ProviderEvent::TextChunk { @@ -246,7 +288,7 @@ fn map_sdk_event_to_provider_event( } Some(ProviderEvent::Thinking { thinking: text, - timestamp: chrono::Utc::now().naive_utc(), + timestamp: now, }) } Part::Tool(tool_part) => match &tool_part.state { @@ -255,7 +297,7 @@ fn map_sdk_event_to_provider_event( tool_use_id: Some(tool_part.call_id.clone()), input: tool_part.input.clone().unwrap_or(serde_json::Value::Null), message_id: data.message_id.clone(), - timestamp: chrono::Utc::now().naive_utc(), + timestamp: now, }), ToolState::Completed(state) => { let output = state.output.trim().to_string(); @@ -266,41 +308,38 @@ fn map_sdk_event_to_provider_event( tool_use_id: Some(tool_part.call_id.clone()), result: output, message_id: data.message_id.clone(), - timestamp: chrono::Utc::now().naive_utc(), + timestamp: now, }) } ToolState::Error(state) => Some(ProviderEvent::Error { error: state.error.clone(), - timestamp: chrono::Utc::now().naive_utc(), + timestamp: now, }), ToolState::Pending(_) => None, }, _ => None, }, Event::SessionStatus(data) => { - if data.session_id == session_id { - if data.status.status_type == "error" { - Some(ProviderEvent::Error { - error: "session error".into(), - timestamp: chrono::Utc::now().naive_utc(), - }) - } else if data.status.status_type == "idle" { - Some(ProviderEvent::Done { - data: serde_json::json!({}), - timestamp: chrono::Utc::now().naive_utc(), - }) - } else { - None - } - } else { - None + if data.session_id != session_id { + return None; + } + match data.status.status_type.as_str() { + "error" => Some(ProviderEvent::Error { + error: "session error".into(), + timestamp: now, + }), + "idle" => Some(ProviderEvent::Done { + data: serde_json::json!({}), + timestamp: now, + }), + _ => None, } } Event::SessionIdle(data) => { if data.session_id == session_id { Some(ProviderEvent::Done { data: serde_json::json!({}), - timestamp: chrono::Utc::now().naive_utc(), + timestamp: now, }) } else { None @@ -310,7 +349,7 @@ fn map_sdk_event_to_provider_event( if data.session_id == session_id { Some(ProviderEvent::Error { error: data.error_message(), - timestamp: chrono::Utc::now().naive_utc(), + timestamp: now, }) } else { None @@ -343,7 +382,7 @@ fn map_sdk_event_to_provider_event( .collect(), tool_call_id: None, message_id: None, - timestamp: chrono::Utc::now().naive_utc(), + timestamp: now, }), _ => None, } @@ -398,12 +437,7 @@ impl LlmProvider for OpenCodeSdkProvider { &self, input: TurnInput, ) -> Result, ProviderError> { - let client = self - .client - .lock() - .unwrap() - .clone() - .ok_or(ProviderError::NotStarted)?; + let (client, user_id) = self.client_and_user_id()?; tracing::info!(model = %input.model, "start_turn: creating opencode session"); let session = client @@ -414,12 +448,6 @@ impl LlmProvider for OpenCodeSdkProvider { *self.session_id.lock().unwrap() = Some(session.id.clone()); - let user_id = self - .user_id - .lock() - .unwrap() - .ok_or_else(|| ProviderError::Protocol("user_id not set on provider".into()))?; - // Subscribe to the global event stream BEFORE issuing the prompt so // we don't miss events that fire immediately when the prompt is // queued on the server. @@ -449,12 +477,7 @@ impl LlmProvider for OpenCodeSdkProvider { &self, input: ResumeInput, ) -> Result, ProviderError> { - let client = self - .client - .lock() - .unwrap() - .clone() - .ok_or(ProviderError::NotStarted)?; + let (client, user_id) = self.client_and_user_id()?; // Mirror the reference implementation's `sendTurnMessage` (see // `spec/reference/server/services/providers/opencode/index.ts`): @@ -479,12 +502,6 @@ impl LlmProvider for OpenCodeSdkProvider { .map(|s| s.to_string()) .unwrap_or_else(|| "continue".to_string()); - let user_id = self - .user_id - .lock() - .unwrap() - .ok_or_else(|| ProviderError::Protocol("user_id not set on provider".into()))?; - // Subscribe BEFORE issuing the prompt_async so we don't miss events // that fire immediately when the prompt is queued on the server. let rx = self @@ -511,14 +528,9 @@ impl LlmProvider for OpenCodeSdkProvider { } async fn abort_turn(&self) -> Result<(), ProviderError> { - if let Some(cancellation) = self.event_cancellation.lock().unwrap().take() { - cancellation.cancel(); - } - let (session_id, client) = { - let s = self.session_id.lock().unwrap().clone(); - let c = self.client.lock().unwrap().clone(); - (s, c) - }; + self.cancel_inflight(); + let session_id = self.session_id.lock().unwrap().clone(); + let client = self.client.lock().unwrap().clone(); if let (Some(client), Some(session_id)) = (client, session_id) { let _ = client.session.abort(&session_id).await; } @@ -549,9 +561,7 @@ impl LlmProvider for OpenCodeSdkProvider { // it is reaped by the idle-reaper task or by the process-exit // handlers in `src/main.rs` (which call // `OpenCodeServerPool::instance().shutdown_all()`). - if let Some(cancellation) = self.event_cancellation.lock().unwrap().take() { - cancellation.cancel(); - } + self.cancel_inflight(); *self.client.lock().unwrap() = None; Ok(true) } @@ -561,6 +571,28 @@ impl LlmProvider for OpenCodeSdkProvider { mod tests { use super::*; + fn test_provider(snippet: &str) -> OpenCodeSdkProvider { + OpenCodeSdkProvider { + config: HarnessConfig { + agent_type: "test".into(), + harness: "opencode".into(), + provider_config_ref: "test.json".into(), + model: None, + effort: None, + scope: crate::db::schema::ScopeType::Global, + }, + provider_snippet: snippet.into(), + config_root: PathBuf::from("/tmp"), + client: Mutex::new(None), + session_id: Mutex::new(None), + event_cancellation: Mutex::new(None), + reader_cancellation: Mutex::new(None), + user_id: Mutex::new(None), + working_dir: Mutex::new(None), + log_data: false, + } + } + #[test] fn test_event_mapping_text_chunk() { let global = GlobalEvent { @@ -731,24 +763,7 @@ mod tests { #[test] fn test_extract_provider_id() { let snippet = r#"{"provider": {"anthropic": {"apiKey": "sk-..."}}}"#; - let provider = OpenCodeSdkProvider { - config: HarnessConfig { - agent_type: "test".into(), - harness: "opencode".into(), - provider_config_ref: "test.json".into(), - model: None, - effort: None, - scope: crate::db::schema::ScopeType::Global, - }, - provider_snippet: snippet.into(), - config_root: PathBuf::from("/tmp"), - client: Mutex::new(None), - session_id: Mutex::new(None), - event_cancellation: Mutex::new(None), - user_id: Mutex::new(None), - working_dir: Mutex::new(None), - log_data: false, - }; + let provider = test_provider(snippet); assert_eq!(provider.extract_provider_id(), Some("anthropic".into())); } @@ -776,24 +791,7 @@ mod tests { #[test] fn test_extract_provider_id_empty() { - let provider = OpenCodeSdkProvider { - config: HarnessConfig { - agent_type: "test".into(), - harness: "opencode".into(), - provider_config_ref: "test.json".into(), - model: None, - effort: None, - scope: crate::db::schema::ScopeType::Global, - }, - provider_snippet: "{}".into(), - config_root: PathBuf::from("/tmp"), - client: Mutex::new(None), - session_id: Mutex::new(None), - event_cancellation: Mutex::new(None), - user_id: Mutex::new(None), - working_dir: Mutex::new(None), - log_data: false, - }; + let provider = test_provider("{}"); assert_eq!(provider.extract_provider_id(), None); } } diff --git a/src/server/ws/message.rs b/src/server/ws/message.rs index e6a7976..4f591b8 100644 --- a/src/server/ws/message.rs +++ b/src/server/ws/message.rs @@ -87,7 +87,7 @@ mod tests { }; let json = serde_json::to_string(&msg).unwrap(); let deserialized: ClientMessage = serde_json::from_str(&json).unwrap(); - assert_eq!(json.contains("\"type\":\"subscribe\""), true); + assert!(json.contains("\"type\":\"subscribe\"")); match deserialized { ClientMessage::Subscribe { topics, since } => { assert_eq!(topics.len(), 1); diff --git a/src/services/transcript.rs b/src/services/transcript.rs index 5a84bf9..7047584 100644 --- a/src/services/transcript.rs +++ b/src/services/transcript.rs @@ -9,7 +9,6 @@ pub async fn persist_event( session_id: &str, project_key: i64, ) -> Result<(), hiqlite::Error> { - let seq: i32 = next_seq(client, session_id, project_key).await?; let entry_json = serde_json::to_value(event) .map_err(|e| hiqlite::Error::new(format!("serialize event: {e}")))?; @@ -19,8 +18,11 @@ pub async fn persist_event( client .execute( - "INSERT INTO messages (project_key, session_id, seq, entry_json, timestamp) VALUES ($1, $2, $3, $4, $5)", - hiqlite::params!(project_key, session_id, seq, entry_json.to_string(), timestamp), + "INSERT INTO messages (project_key, session_id, seq, entry_json, timestamp) + VALUES ($1, $2, + (SELECT COALESCE(MAX(seq), 0) + 1 FROM messages WHERE project_key = $1 AND session_id = $2), + $3, $4)", + hiqlite::params!(project_key, session_id, entry_json.to_string(), timestamp), ) .await?; Ok(()) @@ -39,24 +41,6 @@ fn event_timestamp(event: &ProviderEvent) -> Option { } } -async fn next_seq( - client: &Client, - session_id: &str, - project_key: i64, -) -> Result { - let mut rows = client - .query_raw( - "SELECT COALESCE(MAX(seq), 0) + 1 AS next_seq FROM messages WHERE project_key = $1 AND session_id = $2", - hiqlite::params!(project_key, session_id), - ) - .await?; - let seq: i64 = rows - .first_mut() - .expect("COALESCE query always returns one row") - .get("next_seq"); - Ok(seq as i32) -} - pub async fn load_transcript( client: &Client, session_id: &str, @@ -224,6 +208,54 @@ mod tests { assert_eq!(messages[2].seq, 3); } + #[tokio::test] + async fn test_persist_event_concurrent() { + let (client, _tmp) = make_client().await; + let session_id = "sess-concurrent"; + let project_key = 4i64; + let concurrency = 20; + + let ts = NaiveDateTime::parse_from_str("2024-01-15 12:00:00", "%Y-%m-%d %H:%M:%S").unwrap(); + let mut handles = Vec::new(); + for i in 0..concurrency { + let event = ProviderEvent::Text { + text: format!("event-{i}"), + timestamp: ts, + }; + let cl = client.clone(); + let sid = session_id.to_string(); + handles.push(tokio::spawn(async move { + persist_event(&cl, &event, &sid, project_key).await + })); + } + + for handle in handles { + handle.await.unwrap().unwrap(); + } + + let messages = client + .query_map::( + "SELECT project_key, session_id, seq, entry_json FROM messages WHERE project_key = $1 AND session_id = $2 ORDER BY seq ASC", + hiqlite::params!(project_key, session_id), + ) + .await + .unwrap(); + + assert_eq!(messages.len(), concurrency); + let seqs: Vec = messages.iter().map(|m| m.seq).collect(); + let mut sorted = seqs.clone(); + sorted.sort(); + assert_eq!( + seqs, sorted, + "seq values should be monotonically increasing" + ); + // Verify all seq values are unique and start at 1 + let unique: std::collections::HashSet = seqs.iter().copied().collect(); + assert_eq!(unique.len(), concurrency); + assert_eq!(seqs[0], 1); + assert_eq!(seqs[concurrency - 1], concurrency as i32); + } + #[tokio::test] async fn test_persist_multiple_sessions() { let (client, _tmp) = make_client().await; diff --git a/tests/agent_runs_test.rs b/tests/agent_runs_test.rs index be9a6c3..b23e116 100644 --- a/tests/agent_runs_test.rs +++ b/tests/agent_runs_test.rs @@ -175,7 +175,7 @@ async fn test_create_agent_run_201() { assert_eq!(body["status"], "running"); assert_eq!(body["task_id"].as_i64().unwrap(), task_id); assert_eq!(body["agent_type"], "implementation"); - assert!(body["id"].as_str().unwrap().len() > 0); + assert!(!body["id"].as_str().unwrap().is_empty()); } #[tokio::test] diff --git a/tests/hiqlite_ports_test.rs b/tests/hiqlite_ports_test.rs index d2a5b50..318cd04 100644 --- a/tests/hiqlite_ports_test.rs +++ b/tests/hiqlite_ports_test.rs @@ -106,18 +106,14 @@ async fn test_server_with_env_configured_hiqlite_ports() { // Health endpoint let resp = client - .get(&format!("http://{}/health", addr)) + .get(format!("http://{}/health", addr)) .send() .await .unwrap(); assert_eq!(resp.status(), 200); // GET / should redirect to /webapp - let resp = client - .get(&format!("http://{}", addr)) - .send() - .await - .unwrap(); + let resp = client.get(format!("http://{}", addr)).send().await.unwrap(); assert!( resp.status().is_redirection(), "expected redirect, got {}", @@ -126,7 +122,7 @@ async fn test_server_with_env_configured_hiqlite_ports() { // GET /webapp should serve the shell page let resp = client - .get(&format!("http://{}/webapp", addr)) + .get(format!("http://{}/webapp", addr)) .send() .await .unwrap(); @@ -152,7 +148,7 @@ async fn test_hiqlite_ports_do_not_use_zero() { let client = reqwest::Client::new(); let resp = client - .get(&format!("http://{}/health", addr)) + .get(format!("http://{}/health", addr)) .send() .await .unwrap(); diff --git a/tests/opencode_sdk_integration_test.rs b/tests/opencode_sdk_integration_test.rs index 1a84823..558440b 100644 --- a/tests/opencode_sdk_integration_test.rs +++ b/tests/opencode_sdk_integration_test.rs @@ -123,9 +123,8 @@ async fn test_opencode_sdk_server_shutdown_releases_port() { ) .await; - match probe { - Ok(Ok(_)) => panic!("port {port} should be free after shutdown"), - _ => {} // connection refused or timed out = port is free + if let Ok(Ok(_)) = probe { + panic!("port {port} should be free after shutdown"); } } @@ -383,7 +382,7 @@ async fn test_opencode_sdk_unstructured_conversation() { tracing::info!("unstructured conversation received events: {received}"); let _ = conv.abort().await; - let _ = client.session.delete(&conv.session_id()).await; + let _ = client.session.delete(conv.session_id()).await; server.shutdown().await.unwrap(); } @@ -411,7 +410,7 @@ async fn test_opencode_sdk_session_resume() { tracing::info!("session resume turn 2 received events: {turn2}"); let _ = conv.abort().await; - let _ = client.session.delete(&conv.session_id()).await; + let _ = client.session.delete(conv.session_id()).await; server.shutdown().await.unwrap(); } @@ -435,9 +434,8 @@ async fn test_opencode_sdk_process_leak() { tokio::net::TcpStream::connect(&addr), ) .await; - match probe { - Ok(Ok(_)) => panic!("port {port} should be free after shutdown"), - _ => {} // port is free or timed out + if let Ok(Ok(_)) = probe { + panic!("port {port} should be free after shutdown"); } // Verify no orphan processes — we can't know the exact PID from outside @@ -483,7 +481,7 @@ async fn test_opencode_sdk_multi_session_lifecycle() { // Clean up all conversations for conv in &conversations { let _ = conv.abort().await; - let _ = client.session.delete(&conv.session_id()).await; + let _ = client.session.delete(conv.session_id()).await; } server.shutdown().await.unwrap(); diff --git a/tests/projects_test.rs b/tests/projects_test.rs index 63dde9f..c20668a 100644 --- a/tests/projects_test.rs +++ b/tests/projects_test.rs @@ -143,7 +143,7 @@ async fn test_list_projects() { assert_eq!(resp.status(), 200); let body: serde_json::Value = resp.json().await.unwrap(); - assert!(body.as_array().unwrap().len() >= 1); + assert!(!body.as_array().unwrap().is_empty()); assert_eq!(body[0]["name"], "test-project"); }