Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
149 changes: 117 additions & 32 deletions apps/desktop/src-tauri/src/lib.rs
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
use std::collections::{HashMap, VecDeque};
use std::error::Error;
use std::sync::{Arc, Mutex};

Expand Down Expand Up @@ -28,6 +29,7 @@ const DEFAULT_GATEWAY_ADDRESS: &str = "127.0.0.1:8787";
const DEFAULT_OPENAI_UPSTREAM: &str = "https://api.openai.com";
const CAPTURE_UPDATED_EVENT: &str = "capture-updated";
const CAPTURE_RUNTIME_ERROR_EVENT: &str = "capture-runtime-error";
const MAX_PENDING_CAPTURE_OUTCOMES: usize = 256;

type SharedStore = Arc<Mutex<Option<EncryptedStore>>>;

Expand Down Expand Up @@ -233,51 +235,92 @@ async fn process_capture_events(
}
};
let adapters = AdapterRegistry::default();
let mut pending = HashMap::<String, GatewayCaptureEvent>::new();

while let Some(event) = receiver.recv().await {
let result = {
let mut store = match store.lock() {
Ok(store) => store,
Err(_) => {
let mut ready = VecDeque::from([event]);
while let Some(event) = ready.pop_front() {
let retry_event = event.clone();
let result = {
let mut store = match store.lock() {
Ok(store) => store,
Err(_) => {
emit_runtime_error(
&app,
"capture_store_unavailable",
"encrypted workspace is temporarily unavailable".to_owned(),
);
return;
}
};
let Some(store) = store.as_mut() else {
emit_runtime_error(
&app,
"capture_store_unavailable",
"encrypted workspace is temporarily unavailable".to_owned(),
"capture_store_uninitialized",
"encrypted workspace has not initialized".to_owned(),
);
return;
}
};
process_gateway_event(store, &policy, &adapters, event)
};
let Some(store) = store.as_mut() else {
emit_runtime_error(
&app,
"capture_store_uninitialized",
"encrypted workspace has not initialized".to_owned(),
);
return;
};
process_gateway_event(store, &policy, &adapters, event)
};

match result {
Ok(GatewayCaptureOutcome::Persisted(capture)) => {
let _ = app.emit(
CAPTURE_UPDATED_EVENT,
CaptureUpdated {
capture_id: capture.capture_id,
},
);
}
Ok(
GatewayCaptureOutcome::ResponseObserved(_)
| GatewayCaptureOutcome::UpstreamFailed(_),
) => {}
Err(error) => {
emit_runtime_error(&app, "capture_processing_failed", error.to_string());
match result {
Ok(GatewayCaptureOutcome::Persisted(capture)) => {
emit_capture_updated(&app, capture.capture_id.clone());
if let Some(outcome) = pending.remove(&capture.capture_id) {
ready.push_back(outcome);
}
}
Ok(GatewayCaptureOutcome::ResponseObserved(response)) => {
if response.persisted {
emit_capture_updated(&app, response.capture_id);
} else if !queue_pending_outcome(&mut pending, retry_event) {
emit_runtime_error(
&app,
"capture_pending_overflow",
"too many capture outcomes arrived before their requests".to_owned(),
);
}
}
Ok(GatewayCaptureOutcome::UpstreamFailed(failure)) => {
if failure.persisted {
emit_capture_updated(&app, failure.capture_id);
} else if !queue_pending_outcome(&mut pending, retry_event) {
emit_runtime_error(
&app,
"capture_pending_overflow",
"too many capture outcomes arrived before their requests".to_owned(),
);
}
}
Err(error) => {
emit_runtime_error(&app, "capture_processing_failed", error.to_string());
}
}
}
}
}

fn queue_pending_outcome(
pending: &mut HashMap<String, GatewayCaptureEvent>,
event: GatewayCaptureEvent,
) -> bool {
let capture_id = match &event {
GatewayCaptureEvent::Request(_) => return false,
GatewayCaptureEvent::Response(response) => response.capture_id.clone(),
GatewayCaptureEvent::UpstreamFailure(failure) => failure.capture_id.clone(),
};
if pending.len() >= MAX_PENDING_CAPTURE_OUTCOMES && !pending.contains_key(&capture_id) {
return false;
}
pending.insert(capture_id, event);
true
}

fn emit_capture_updated(app: &AppHandle, capture_id: String) {
let _ = app.emit(CAPTURE_UPDATED_EVENT, CaptureUpdated { capture_id });
}

fn apply_gateway_state(workspace: &mut WorkspaceBootstrap, runtime: &GatewayRuntime) {
workspace.capture.active = runtime.capture.is_enabled();
workspace.capture.can_control = true;
Expand Down Expand Up @@ -315,6 +358,7 @@ fn show_main_window(app: &tauri::AppHandle) {

#[cfg(test)]
mod tests {
use codeischeap_gateway::{CapturedPayload, GatewayResponseCapture};
use codeischeap_storage::DatabaseKey;
use tempfile::tempdir;

Expand Down Expand Up @@ -400,4 +444,45 @@ mod tests {
assert_eq!(workspace.capture.endpoint, "http://127.0.0.1:8787");
assert_eq!(workspace.capture.profile, "OpenAI-compatible local gateway");
}

#[test]
fn pending_outcomes_are_keyed_replaced_and_bounded() {
let mut pending = HashMap::new();
for index in 0..MAX_PENDING_CAPTURE_OUTCOMES {
assert!(queue_pending_outcome(
&mut pending,
response_event(&format!("capture_{index}"), 200)
));
}
assert_eq!(pending.len(), MAX_PENDING_CAPTURE_OUTCOMES);
assert!(queue_pending_outcome(
&mut pending,
response_event("capture_0", 429)
));
assert!(!queue_pending_outcome(
&mut pending,
response_event("overflow", 200)
));
let GatewayCaptureEvent::Response(replaced) = pending
.get("capture_0")
.expect("existing outcome must remain")
else {
panic!("pending event must be a response");
};
assert_eq!(replaced.status, 429);
}

fn response_event(capture_id: &str, status: u16) -> GatewayCaptureEvent {
GatewayCaptureEvent::Response(GatewayResponseCapture {
capture_id: capture_id.to_owned(),
status,
headers: Vec::new(),
duration_ms: 1,
body: CapturedPayload {
bytes: Vec::new().into(),
truncated: false,
complete: true,
},
})
}
}
1 change: 1 addition & 0 deletions crates/adapters/tests/anthropic.rs
Original file line number Diff line number Diff line change
Expand Up @@ -167,6 +167,7 @@ fn sanitized_anthropic(capture_id: &str, path: &str, body: CapturedBody) -> Sani
headers: Vec::new(),
body,
},
outcome: None,
redactions: Vec::new(),
};
CapturePolicy::load_default()
Expand Down
1 change: 1 addition & 0 deletions crates/adapters/tests/openai.rs
Original file line number Diff line number Diff line change
Expand Up @@ -291,6 +291,7 @@ fn sanitized(capture_id: &str, host: &str, path: &str, body: CapturedBody) -> Sa
headers: Vec::new(),
body,
},
outcome: None,
redactions: Vec::new(),
};
CapturePolicy::load_default()
Expand Down
37 changes: 37 additions & 0 deletions crates/capture-ipc/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -25,6 +25,8 @@ pub struct CaptureEnvelope {
pub observed_at_unix_ms: u64,
pub source: CaptureSource,
pub request: CapturedRequest,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub outcome: Option<CaptureOutcome>,
#[serde(default)]
pub redactions: Vec<CaptureRedaction>,
}
Expand Down Expand Up @@ -70,12 +72,45 @@ pub struct CapturedBody {
pub enum CapturedBodyState {
Empty,
Json,
Text,
InvalidJson,
InvalidUtf8,
Truncated,
OmittedUnsupportedContentType,
}

#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, JsonSchema)]
#[serde(tag = "kind", content = "result", rename_all = "snake_case")]
pub enum CaptureOutcome {
Response(CapturedResponse),
UpstreamFailure(CapturedUpstreamFailure),
}

#[derive(Debug, Clone, PartialEq, Serialize, Deserialize, JsonSchema)]
#[serde(deny_unknown_fields)]
pub struct CapturedResponse {
pub status: u16,
#[serde(default)]
pub headers: Vec<CapturedField>,
pub body: CapturedBody,
pub duration_ms: u64,
pub completeness: ResponseCompleteness,
}

#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
#[serde(deny_unknown_fields)]
pub struct CapturedUpstreamFailure {
pub duration_ms: u64,
}

#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
#[serde(rename_all = "snake_case")]
pub enum ResponseCompleteness {
Complete,
Truncated,
Incomplete,
}

#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize, JsonSchema)]
#[serde(deny_unknown_fields)]
pub struct CaptureRedaction {
Expand All @@ -89,6 +124,8 @@ pub enum RedactionLocation {
Header,
Query,
Body,
ResponseHeader,
ResponseBody,
}

#[derive(Debug, Deserialize)]
Expand Down
33 changes: 31 additions & 2 deletions crates/capture-ipc/tests/contract.rs
Original file line number Diff line number Diff line change
@@ -1,6 +1,7 @@
use codeischeap_capture_ipc::{
CAPTURE_ENVELOPE_VERSION, CaptureEnvelope, CaptureSource, CapturedBody, CapturedBodyState,
CapturedRequest, IpcError, receive_from_reader,
CAPTURE_ENVELOPE_VERSION, CaptureEnvelope, CaptureOutcome, CaptureSource, CapturedBody,
CapturedBodyState, CapturedField, CapturedRequest, CapturedResponse, IpcError,
ResponseCompleteness, receive_from_reader,
};
use schemars::schema_for;
use tokio::io::{AsyncWriteExt, BufReader};
Expand Down Expand Up @@ -29,6 +30,7 @@ fn sample_envelope() -> CaptureEnvelope {
),
},
},
outcome: None,
redactions: Vec::new(),
}
}
Expand Down Expand Up @@ -68,6 +70,33 @@ async fn accepts_an_authenticated_envelope() {
assert_eq!(received, expected);
}

#[tokio::test]
async fn response_outcomes_round_trip_through_authenticated_ipc() {
let mut expected = sample_envelope();
expected.outcome = Some(CaptureOutcome::Response(CapturedResponse {
status: 200,
headers: vec![CapturedField {
name: "content-type".to_owned(),
value: "text/event-stream".to_owned(),
}],
body: CapturedBody {
state: CapturedBodyState::Text,
content: Some(serde_json::Value::String(
"data: {\"type\":\"done\"}\n\n".to_owned(),
)),
},
duration_ms: 73,
completeness: ResponseCompleteness::Complete,
}));
let mut reader = framed_reader("synthetic-token", &expected).await;

let received = receive_from_reader(&mut reader, "synthetic-token")
.await
.expect("response outcome must be accepted");

assert_eq!(received, expected);
}

#[tokio::test]
async fn rejects_an_invalid_token_without_echoing_it() {
let mut reader = framed_reader("wrong-token", &sample_envelope()).await;
Expand Down
32 changes: 25 additions & 7 deletions crates/capture-policy/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -4,7 +4,8 @@ use std::collections::HashSet;
use std::fmt;

use codeischeap_capture_ipc::{
CaptureEnvelope, CaptureRedaction, CapturedField, CapturedRequest, RedactionLocation,
CaptureEnvelope, CaptureOutcome, CaptureRedaction, CapturedField, CapturedRequest,
RedactionLocation,
};
use schemars::JsonSchema;
use serde::{Deserialize, Serialize};
Expand Down Expand Up @@ -191,11 +192,30 @@ impl CapturePolicy {
if let Some(content) = envelope.request.body.content.as_mut() {
scrub_json(
content,
RedactionLocation::Body,
&sensitive_names,
&mut envelope.redactions,
&mut newly_redacted,
);
}
if let Some(CaptureOutcome::Response(response)) = envelope.outcome.as_mut() {
scrub_fields(
&mut response.headers,
RedactionLocation::ResponseHeader,
&sensitive_names,
&mut envelope.redactions,
&mut newly_redacted,
);
if let Some(content) = response.body.content.as_mut() {
scrub_json(
content,
RedactionLocation::ResponseBody,
&sensitive_names,
&mut envelope.redactions,
&mut newly_redacted,
);
}
}

Ok(SanitizedCapture {
envelope,
Expand Down Expand Up @@ -235,6 +255,7 @@ fn scrub_fields(

fn scrub_json(
value: &mut serde_json::Value,
location: RedactionLocation,
sensitive_names: &HashSet<String>,
redactions: &mut Vec<CaptureRedaction>,
newly_redacted: &mut usize,
Expand All @@ -248,19 +269,16 @@ fn scrub_json(
.collect();
for name in removed {
object.remove(&name);
redactions.push(CaptureRedaction {
location: RedactionLocation::Body,
name,
});
redactions.push(CaptureRedaction { location, name });
*newly_redacted += 1;
}
for child in object.values_mut() {
scrub_json(child, sensitive_names, redactions, newly_redacted);
scrub_json(child, location, sensitive_names, redactions, newly_redacted);
}
}
serde_json::Value::Array(items) => {
for item in items {
scrub_json(item, sensitive_names, redactions, newly_redacted);
scrub_json(item, location, sensitive_names, redactions, newly_redacted);
}
}
_ => {}
Expand Down
Loading
Loading