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
50 changes: 1 addition & 49 deletions data_connector/src/common.rs
Original file line number Diff line number Diff line change
@@ -1,5 +1,3 @@
use std::collections::HashMap;

use serde_json::Value;

use crate::{
Expand All @@ -18,10 +16,6 @@ pub(super) const RESPONSE_COLUMNS: &[&str] = &[
"conversation_id",
"previous_response_id",
"input",
"instructions",
"output",
"tool_calls",
"metadata",
"created_at",
"safety_identifier",
"model",
Expand Down Expand Up @@ -131,20 +125,6 @@ pub(super) fn parse_conversation_metadata(
}
}

pub(super) fn parse_tool_calls(raw: Option<String>) -> Result<Vec<Value>, String> {
match raw {
Some(s) if !s.is_empty() => serde_json::from_str(&s).map_err(|e| e.to_string()),
_ => Ok(Vec::new()),
}
}

pub(super) fn parse_metadata(raw: Option<String>) -> Result<HashMap<String, Value>, String> {
match raw {
Some(s) if !s.is_empty() => serde_json::from_str(&s).map_err(|e| e.to_string()),
_ => Ok(HashMap::new()),
}
}

pub(super) fn parse_raw_response(raw: Option<String>) -> Result<Value, String> {
match raw {
Some(s) if !s.is_empty() => serde_json::from_str(&s).map_err(|e| e.to_string()),
Expand All @@ -165,34 +145,6 @@ mod tests {

use super::*;

#[test]
fn parse_tool_calls_handles_empty_input() {
assert!(parse_tool_calls(None).unwrap().is_empty());
assert!(parse_tool_calls(Some(String::new())).unwrap().is_empty());
}

#[test]
fn parse_tool_calls_round_trips() {
let payload = json!([{ "type": "test", "value": 1 }]).to_string();
let parsed = parse_tool_calls(Some(payload)).unwrap();
assert_eq!(parsed.len(), 1);
assert_eq!(parsed[0]["type"], "test");
assert_eq!(parsed[0]["value"], 1);
}

#[test]
fn parse_metadata_defaults_to_empty_map() {
assert!(parse_metadata(None).unwrap().is_empty());
}

#[test]
fn parse_metadata_round_trips() {
let payload = json!({"key": "value", "nested": {"bool": true}}).to_string();
let parsed = parse_metadata(Some(payload)).unwrap();
assert_eq!(parsed.get("key").unwrap(), "value");
assert_eq!(parsed["nested"]["bool"], true);
}

#[test]
fn parse_raw_response_handles_null() {
assert_eq!(parse_raw_response(None).unwrap(), Value::Null);
Expand Down Expand Up @@ -588,7 +540,7 @@ mod tests {
);

// Core columns that are not skipped must still be present
for col in &["id", "input", "output", "model", "conversation_id"] {
for col in &["id", "input", "model", "conversation_id"] {
assert!(
sql.contains(col),
"core column '{col}' should remain: {sql}"
Expand Down
41 changes: 18 additions & 23 deletions data_connector/src/core.rs
Original file line number Diff line number Diff line change
Expand Up @@ -347,18 +347,6 @@ pub struct StoredResponse {
/// Input items as JSON array
pub input: Value,

/// System instructions used
pub instructions: Option<String>,

/// Output items as JSON array
pub output: Value,

/// Tool calls made by the model (if any)
pub tool_calls: Vec<Value>,

/// Custom metadata
pub metadata: HashMap<String, Value>,

/// When this response was created
pub created_at: DateTime<Utc>,

Expand All @@ -383,10 +371,6 @@ impl StoredResponse {
id: ResponseId::new(),
previous_response_id,
input: Value::Array(vec![]),
instructions: None,
output: Value::Array(vec![]),
tool_calls: Vec::new(),
metadata: HashMap::new(),
created_at: Utc::now(),
safety_identifier: None,
model: None,
Expand Down Expand Up @@ -441,7 +425,14 @@ impl ResponseChain {

responses
.iter()
.map(|r| (r.input.clone(), r.output.clone()))
.map(|r| {
let output = r
.raw_response
.get("output")
.cloned()
.unwrap_or(Value::Array(vec![]));
(r.input.clone(), output)
})
.collect()
}
}
Expand Down Expand Up @@ -865,9 +856,9 @@ mod tests {
"default input should be empty array"
);
assert_eq!(
resp.output,
Value::Array(vec![]),
"default output should be empty array"
resp.raw_response,
Value::Null,
"default raw_response should be Null"
);
}

Expand Down Expand Up @@ -934,15 +925,17 @@ mod tests {

#[test]
fn response_chain_build_context_returns_input_output_pairs() {
use serde_json::json;

let mut chain = ResponseChain::new();

let mut r1 = StoredResponse::new(None);
r1.input = Value::String("input1".to_string());
r1.output = Value::String("output1".to_string());
r1.raw_response = json!({"output": "output1"});

let mut r2 = StoredResponse::new(None);
r2.input = Value::String("input2".to_string());
r2.output = Value::String("output2".to_string());
r2.raw_response = json!({"output": "output2"});

chain.add_response(r1);
chain.add_response(r2);
Expand All @@ -957,12 +950,14 @@ mod tests {

#[test]
fn response_chain_build_context_with_max_responses_limits_output() {
use serde_json::json;

let mut chain = ResponseChain::new();

for i in 0..5 {
let mut resp = StoredResponse::new(None);
resp.input = Value::String(format!("input{i}"));
resp.output = Value::String(format!("output{i}"));
resp.raw_response = json!({"output": format!("output{i}")});
chain.add_response(resp);
}

Expand Down
4 changes: 2 additions & 2 deletions data_connector/src/hooked.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1252,7 +1252,7 @@ mod tests {

let mut resp = StoredResponse::new(None);
resp.input = json!("round-trip-input");
resp.output = json!(["round-trip-output"]);
resp.raw_response = json!({"output": ["round-trip-output"]});
resp.safety_identifier = Some("user-rt".to_string());

let id = hooked.store_response(resp).await.unwrap();
Expand All @@ -1263,7 +1263,7 @@ mod tests {
assert!(fetched.is_some(), "stored response should be retrievable");
let fetched = fetched.unwrap();
assert_eq!(fetched.input, json!("round-trip-input"));
assert_eq!(fetched.output, json!(["round-trip-output"]));
assert_eq!(fetched.raw_response["output"], json!(["round-trip-output"]));
assert_eq!(fetched.safety_identifier.as_deref(), Some("user-rt"));

// get_response also triggers before/after hooks
Expand Down
17 changes: 6 additions & 11 deletions data_connector/src/memory.rs
Original file line number Diff line number Diff line change
Expand Up @@ -600,14 +600,14 @@ mod tests {
let mut response = StoredResponse::new(None);
response.id = ResponseId::from("resp_custom");
response.input = json!("Input");
response.output = json!("Output");
response.raw_response = json!({"output": "Output"});
store.store_response(response.clone()).await.unwrap();
let retrieved = store
.get_response(&ResponseId::from("resp_custom"))
.await
.unwrap();
assert!(retrieved.is_some());
assert_eq!(retrieved.unwrap().output, json!("Output"));
assert_eq!(retrieved.unwrap().raw_response["output"], json!("Output"));
}

#[tokio::test]
Expand All @@ -617,7 +617,7 @@ mod tests {
// Store a response
let mut response = StoredResponse::new(None);
response.input = json!("Hello");
response.output = json!("Hi there!");
response.raw_response = json!({"output": "Hi there!"});
let response_id = store.store_response(response).await.unwrap();

// Retrieve it
Expand All @@ -638,17 +638,17 @@ mod tests {
// Create a chain of responses
let mut response1 = StoredResponse::new(None);
response1.input = json!("First");
response1.output = json!("First response");
response1.raw_response = json!({"output": "First response"});
let id1 = store.store_response(response1).await.unwrap();

let mut response2 = StoredResponse::new(Some(id1.clone()));
response2.input = json!("Second");
response2.output = json!("Second response");
response2.raw_response = json!({"output": "Second response"});
let id2 = store.store_response(response2).await.unwrap();

let mut response3 = StoredResponse::new(Some(id2.clone()));
response3.input = json!("Third");
response3.output = json!("Third response");
response3.raw_response = json!({"output": "Third response"});
let id3 = store.store_response(response3).await.unwrap();

// Get the chain
Expand All @@ -671,19 +671,16 @@ mod tests {
// Store responses for different users
let mut response1 = StoredResponse::new(None);
response1.input = json!("User1 message");
response1.output = json!("Response to user1");
response1.safety_identifier = Some("user1".to_string());
store.store_response(response1).await.unwrap();

let mut response2 = StoredResponse::new(None);
response2.input = json!("Another user1 message");
response2.output = json!("Another response to user1");
response2.safety_identifier = Some("user1".to_string());
store.store_response(response2).await.unwrap();

let mut response3 = StoredResponse::new(None);
response3.input = json!("User2 message");
response3.output = json!("Response to user2");
response3.safety_identifier = Some("user2".to_string());
store.store_response(response3).await.unwrap();

Expand Down Expand Up @@ -725,13 +722,11 @@ mod tests {

let mut response1 = StoredResponse::new(None);
response1.input = json!("Test1");
response1.output = json!("Reply1");
response1.safety_identifier = Some("user1".to_string());
store.store_response(response1).await.unwrap();

let mut response2 = StoredResponse::new(None);
response2.input = json!("Test2");
response2.output = json!("Reply2");
response2.safety_identifier = Some("user2".to_string());
store.store_response(response2).await.unwrap();

Expand Down
48 changes: 3 additions & 45 deletions data_connector/src/oracle.rs
Original file line number Diff line number Diff line change
Expand Up @@ -26,8 +26,8 @@ use super::core::{
};
use crate::{
common::{
build_response_select_base, extra_column_defs, parse_json_value, parse_metadata,
parse_raw_response, parse_tool_calls, resolve_extra_column_values,
build_response_select_base, extra_column_defs, parse_json_value, parse_raw_response,
resolve_extra_column_values,
},
config::OracleConfig,
context::current_extra_columns,
Expand Down Expand Up @@ -1216,14 +1216,10 @@ impl OracleResponseStorage {

if exists == 0 {
let mut col_defs = vec![format!("{} VARCHAR2(64) PRIMARY KEY", s.col("id"))];
let core_cols: [(&str, &str); 11] = [
let core_cols: [(&str, &str); 7] = [
("conversation_id", "VARCHAR2(64)"),
("previous_response_id", "VARCHAR2(64)"),
("input", "CLOB"),
("instructions", "CLOB"),
("output", "CLOB"),
("tool_calls", "CLOB"),
("metadata", "CLOB"),
("created_at", "TIMESTAMP WITH TIME ZONE"),
("safety_identifier", "VARCHAR2(128)"),
("model", "VARCHAR2(128)"),
Expand Down Expand Up @@ -1291,26 +1287,6 @@ impl OracleResponseStorage {
} else {
row.get(s.col("input")).map_err(map_oracle_error)?
};
let instructions: Option<String> = if s.is_skipped("instructions") {
None
} else {
row.get(s.col("instructions")).map_err(map_oracle_error)?
};
let output_json: Option<String> = if s.is_skipped("output") {
None
} else {
row.get(s.col("output")).map_err(map_oracle_error)?
};
let tool_calls_json: Option<String> = if s.is_skipped("tool_calls") {
None
} else {
row.get(s.col("tool_calls")).map_err(map_oracle_error)?
};
let metadata_json: Option<String> = if s.is_skipped("metadata") {
None
} else {
row.get(s.col("metadata")).map_err(map_oracle_error)?
};
let safety_identifier: Option<String> = if s.is_skipped("safety_identifier") {
None
} else {
Expand All @@ -1335,20 +1311,13 @@ impl OracleResponseStorage {
};

let previous_response_id = previous.map(ResponseId);
let tool_calls = parse_tool_calls(tool_calls_json)?;
let metadata = parse_metadata(metadata_json)?;
let raw_response = parse_raw_response(raw_response_json)?;
let input = parse_json_value(input_json)?;
let output = parse_json_value(output_json)?;

Ok(StoredResponse {
id: ResponseId(id),
previous_response_id,
input,
instructions,
output,
tool_calls,
metadata,
created_at,
safety_identifier,
model,
Expand All @@ -1368,10 +1337,6 @@ impl ResponseStorage for OracleResponseStorage {
id,
previous_response_id,
input,
instructions,
output,
tool_calls,
metadata,
created_at,
safety_identifier,
model,
Expand All @@ -1384,9 +1349,6 @@ impl ResponseStorage for OracleResponseStorage {
let response_id_str = id.0;
let previous_id = previous_response_id.map(|r| r.0);
let json_input = serde_json::to_string(&input)?;
let json_output = serde_json::to_string(&output)?;
let json_tool_calls = serde_json::to_string(&tool_calls)?;
let json_metadata = serde_json::to_string(&metadata)?;
let json_raw_response = serde_json::to_string(&raw_response)?;
let schema = self.store.schema.clone();
// Capture extra columns before spawn_blocking (task-locals don't propagate)
Expand All @@ -1402,10 +1364,6 @@ impl ResponseStorage for OracleResponseStorage {
("id", &response_id_str),
("previous_response_id", &previous_id),
("input", &json_input),
("instructions", &instructions),
("output", &json_output),
("tool_calls", &json_tool_calls),
("metadata", &json_metadata),
("created_at", &created_at),
("safety_identifier", &safety_identifier),
("model", &model),
Expand Down
Loading