Merge branch 'main' into platform-extensions

This commit is contained in:
Douwe Osinga
2025-10-01 16:10:08 -04:00
12 changed files with 308 additions and 101 deletions
+31 -11
View File
@@ -4,9 +4,12 @@ use chrono::{DateTime, Utc};
use futures::stream::{FuturesUnordered, StreamExt};
use futures::{future, FutureExt};
use rmcp::service::ClientInitializeError;
use rmcp::transport::streamable_http_client::StreamableHttpClientTransportConfig;
use rmcp::transport::streamable_http_client::{
AuthRequiredError, StreamableHttpClientTransportConfig, StreamableHttpError,
};
use rmcp::transport::{
ConfigureCommandExt, SseClientTransport, StreamableHttpClientTransport, TokioChildProcess,
ConfigureCommandExt, DynamicTransportError, SseClientTransport, StreamableHttpClientTransport,
TokioChildProcess,
};
use std::collections::HashMap;
use std::process::Stdio;
@@ -209,6 +212,28 @@ async fn child_process_client(
}
}
fn extract_auth_error(
res: &Result<McpClient, ClientInitializeError>,
) -> Option<&AuthRequiredError> {
match res {
Ok(_) => None,
Err(err) => match err {
ClientInitializeError::TransportError {
error: DynamicTransportError { error, .. },
..
} => error
.downcast_ref::<StreamableHttpError<reqwest::Error>>()
.and_then(|auth_error| match auth_error {
StreamableHttpError::AuthRequired(auth_required_error) => {
Some(auth_required_error)
}
_ => None,
}),
_ => None,
},
}
}
impl ExtensionManager {
pub fn new() -> Self {
Self {
@@ -353,15 +378,10 @@ impl ExtensionManager {
),
)
.await;
let client = if let Err(e) = client_res {
// make an attempt at oauth, but failing that, return the original error,
// because this might not have been an auth error at all.
// TODO: when rmcp supports it, we should trigger this flow on 401s with
// WWW-Authenticate headers, not just any init error
let am = match oauth_flow(uri, name).await {
Ok(am) => am,
Err(_) => return Err(e.into()),
};
let client = if let Some(_auth_error) = extract_auth_error(&client_res) {
let am = oauth_flow(uri, name)
.await
.map_err(|_| ExtensionError::SetupError("auth error".to_string()))?;
let client = AuthClient::new(reqwest::Client::default(), am);
let transport = StreamableHttpClientTransport::with_client(
client,
@@ -317,8 +317,8 @@ mod tests {
id: "test_session".to_string(),
working_dir: PathBuf::from(working_dir),
description: "Test session".to_string(),
created_at: "2024-01-01T00:00:00Z".to_string(),
updated_at: "2024-01-01T00:00:00Z".to_string(),
created_at: Default::default(),
updated_at: Default::default(),
schedule_id: Some("test_job".to_string()),
recipe: None,
total_tokens: Some(100),
+2 -15
View File
@@ -379,7 +379,7 @@ impl From<PromptMessage> for Message {
}
}
#[derive(ToSchema, Clone, Copy, PartialEq, Serialize, Deserialize)]
#[derive(ToSchema, Clone, Copy, PartialEq, Serialize, Deserialize, Debug)]
/// Metadata for message visibility
#[serde(rename_all = "camelCase")]
pub struct MessageMetadata {
@@ -462,7 +462,7 @@ fn default_true() -> bool {
true
}
#[derive(ToSchema, Clone, PartialEq, Serialize, Deserialize)]
#[derive(ToSchema, Clone, PartialEq, Serialize, Deserialize, Debug)]
/// A message to or from an LLM
#[serde(rename_all = "camelCase")]
pub struct Message {
@@ -476,19 +476,6 @@ pub struct Message {
pub metadata: MessageMetadata,
}
impl fmt::Debug for Message {
fn fmt(&self, f: &mut fmt::Formatter<'_>) -> fmt::Result {
let joined_content: String = self
.content
.iter()
.map(|c| format!("{c}"))
.collect::<Vec<_>>()
.join(" ");
write!(f, "{:?}: {}", self.role, joined_content)
}
}
fn default_created() -> i64 {
0 // old messages do not have timestamps.
}
+189 -31
View File
@@ -168,20 +168,61 @@ pub fn fix_conversation(conversation: Conversation) -> (Conversation, Vec<String
}
fn fix_messages(messages: Vec<Message>) -> (Vec<Message>, Vec<String>) {
let (messages_1, empty_removed) = remove_empty_messages(messages);
let (messages_2, tool_calling_fixed) = fix_tool_calling(messages_1);
let (messages_3, messages_merged) = merge_consecutive_messages(messages_2);
let (messages_4, lead_trail_fixed) = fix_lead_trail(messages_3);
let (messages_5, populated_if_empty) = populate_if_empty(messages_4);
[
merge_text_content_items,
remove_empty_messages,
fix_tool_calling,
merge_consecutive_messages,
fix_lead_trail,
populate_if_empty,
]
.into_iter()
.fold(
(messages, Vec::new()),
|(msgs, mut all_issues), processor| {
let (new_msgs, issues) = processor(msgs);
all_issues.extend(issues);
(new_msgs, all_issues)
},
)
}
let mut issues = Vec::new();
issues.extend(empty_removed);
issues.extend(tool_calling_fixed);
issues.extend(messages_merged);
issues.extend(lead_trail_fixed);
issues.extend(populated_if_empty);
fn merge_text_content_in_message(mut msg: Message) -> Message {
if msg.role != Role::Assistant {
return msg;
}
msg.content = msg
.content
.into_iter()
.fold(Vec::new(), |mut content, item| {
match item {
MessageContent::Text(text) => {
if let Some(MessageContent::Text(ref mut last)) = content.last_mut() {
last.text.push_str(&text.text);
} else {
content.push(MessageContent::Text(text));
}
}
other => content.push(other),
}
content
});
msg
}
(messages_5, issues)
fn merge_text_content_items(messages: Vec<Message>) -> (Vec<Message>, Vec<String>) {
messages.into_iter().fold(
(Vec::new(), Vec::new()),
|(mut messages, mut issues), message| {
let content_len = message.content.len();
let message = merge_text_content_in_message(message);
if content_len != message.content.len() {
issues.push(String::from("Merged text content"))
}
messages.push(message);
(messages, issues)
},
)
}
fn remove_empty_messages(messages: Vec<Message>) -> (Vec<Message>, Vec<String>) {
@@ -189,7 +230,11 @@ fn remove_empty_messages(messages: Vec<Message>) -> (Vec<Message>, Vec<String>)
let filtered_messages = messages
.into_iter()
.filter(|msg| {
if msg.content.is_empty() {
if msg
.content
.iter()
.all(|c| c.as_text().is_some_and(str::is_empty))
{
issues.push("Removed empty message".to_string());
false
} else {
@@ -402,6 +447,24 @@ mod tests {
use rmcp::model::{CallToolRequestParam, Role};
use rmcp::object;
macro_rules! assert_has_issues_unordered {
($fixed:expr, $issues:expr, $($expected:expr),+ $(,)?) => {
{
let mut expected: Vec<&str> = vec![$($expected),+];
let mut actual: Vec<&str> = $issues.iter().map(|s| s.as_str()).collect();
expected.sort();
actual.sort();
if actual != expected {
panic!(
"assertion failed: issues don't match\nexpected: {:?}\n actual: {:?}. Fixed conversation is:\n{:#?}",
expected, $issues, $fixed,
);
}
}
};
}
fn run_verify(messages: Vec<Message>) -> (Vec<Message>, Vec<String>) {
let (fixed, issues) = fix_conversation(Conversation::new_unvalidated(messages.clone()));
@@ -486,17 +549,15 @@ mod tests {
let (fixed, issues) = run_verify(messages);
assert_eq!(fixed.len(), 3);
assert_eq!(issues.len(), 4);
assert!(issues
.iter()
.any(|i| i.contains("Merged consecutive user messages")));
assert!(issues
.iter()
.any(|i| i.contains("Removed tool response 'orphan_1' from assistant message")));
assert!(issues
.iter()
.any(|i| i.contains("Removed tool request 'bad_req' from user message")));
assert_has_issues_unordered!(
fixed,
issues,
"Merged consecutive assistant messages",
"Merged consecutive user messages",
"Removed tool response 'orphan_1' from assistant message",
"Removed tool request 'bad_req' from user message",
);
assert_eq!(fixed[0].role, Role::User);
assert_eq!(fixed[1].role, Role::Assistant);
@@ -536,10 +597,18 @@ mod tests {
assert_eq!(fixed.len(), 1);
assert!(issues.iter().any(|i| i.contains("Removed empty message")));
assert!(issues
.iter()
.any(|i| i.contains("Removed orphaned tool response 'wrong_id'")));
assert_has_issues_unordered!(
fixed,
issues,
"Removed empty message",
"Removed orphaned tool response 'wrong_id'",
"Removed orphaned tool request 'search_1'",
"Removed orphaned tool request 'search_2'",
"Removed empty message",
"Removed empty message",
"Removed leading assistant message",
"Added placeholder user message to empty conversation",
);
assert_eq!(fixed[0].role, Role::User);
assert_eq!(fixed[0].as_concat_text(), "Hello");
@@ -569,9 +638,12 @@ mod tests {
let (fixed, issues) = fix_conversation(conversation);
assert_eq!(fixed.len(), 5);
assert_eq!(issues.len(), 2);
assert!(issues[0].contains("Removed orphaned tool request"));
assert!(issues[1].contains("Merged consecutive assistant messages"));
assert_has_issues_unordered!(
fixed,
issues,
"Removed orphaned tool request 'toolu_bdrk_018adWbP4X26CfoJU5hkhu3i'",
"Merged consecutive assistant messages"
)
}
#[test]
@@ -592,6 +664,92 @@ mod tests {
];
let (_fixed, issues) = run_verify(messages);
assert_eq!(issues.len(), 0);
assert!(issues.is_empty());
}
#[test]
fn test_merge_text_content_items() {
use crate::conversation::message::MessageContent;
use rmcp::model::{AnnotateAble, RawTextContent};
let mut message = Message::assistant().with_text("Hello");
message.content.push(MessageContent::Text(
RawTextContent {
text: " world".to_string(),
meta: None,
}
.no_annotation(),
));
message.content.push(MessageContent::Text(
RawTextContent {
text: "!".to_string(),
meta: None,
}
.no_annotation(),
));
let messages = vec![
Message::user().with_text("hello"),
message,
Message::user().with_text("thanks"),
];
let (fixed, issues) = run_verify(messages);
assert_eq!(fixed.len(), 3);
assert_has_issues_unordered!(fixed, issues, "Merged text content");
let fixed_msg = &fixed[1];
assert_eq!(fixed_msg.content.len(), 1);
if let MessageContent::Text(text_content) = &fixed_msg.content[0] {
assert_eq!(text_content.text, "Hello world!");
} else {
panic!("Expected text content");
}
}
#[test]
fn test_merge_text_content_items_with_mixed_content() {
use crate::conversation::message::MessageContent;
use rmcp::model::{AnnotateAble, RawTextContent};
let mut image_message = Message::assistant().with_text("Look at");
image_message.content.push(MessageContent::Text(
RawTextContent {
text: " this image:".to_string(),
meta: None,
}
.no_annotation(),
));
image_message = image_message.with_image("", "");
let messages = vec![
Message::user().with_text("hello"),
image_message,
Message::user().with_text("thanks"),
];
let (fixed, issues) = run_verify(messages);
assert_eq!(fixed.len(), 3);
assert_has_issues_unordered!(fixed, issues, "Merged text content");
let fixed_msg = &fixed[1];
assert_eq!(fixed_msg.content.len(), 2);
if let MessageContent::Text(text_content) = &fixed_msg.content[0] {
assert_eq!(text_content.text, "Look at this image:");
} else {
panic!("Expected first item to be text content");
}
if let MessageContent::Image(_) = &fixed_msg.content[1] {
// Good
} else {
panic!("Expected second item to be an image");
}
}
}
+5 -13
View File
@@ -1,7 +1,7 @@
use crate::conversation::Conversation;
use crate::session::Session;
use anyhow::Result;
use chrono::NaiveDateTime;
use chrono::{DateTime, Local, NaiveDateTime, TimeZone, Utc};
use std::fs;
use std::io::{self, BufRead};
use std::path::{Path, PathBuf};
@@ -65,9 +65,9 @@ pub fn load_session(session_name: &str, session_path: &Path) -> Result<Session>
if let Some(obj) = metadata_json.as_object_mut() {
obj.entry("id").or_insert(serde_json::json!(session_name));
obj.entry("created_at")
.or_insert(serde_json::json!(format_timestamp(created_time)?));
.or_insert(serde_json::json!(DateTime::<Utc>::from(created_time)));
obj.entry("updated_at")
.or_insert(serde_json::json!(format_timestamp(modified_time)?));
.or_insert(serde_json::json!(DateTime::<Utc>::from(modified_time)));
obj.entry("extension_data").or_insert(serde_json::json!({}));
obj.entry("message_count").or_insert(serde_json::json!(0));
@@ -97,17 +97,9 @@ pub fn load_session(session_name: &str, session_path: &Path) -> Result<Session>
Ok(session)
}
fn format_timestamp(time: SystemTime) -> Result<String> {
let duration = time.duration_since(std::time::UNIX_EPOCH)?;
let timestamp = chrono::DateTime::from_timestamp(duration.as_secs() as i64, 0)
.unwrap_or_default()
.format("%Y-%m-%d %H:%M:%S")
.to_string();
Ok(timestamp)
}
fn parse_session_timestamp(session_name: &str) -> Option<SystemTime> {
NaiveDateTime::parse_from_str(session_name, "%Y%m%d_%H%M%S")
.ok()
.map(|dt| SystemTime::from(dt.and_utc()))
.and_then(|dt| Local.from_local_datetime(&dt).single())
.map(SystemTime::from)
}
+7 -6
View File
@@ -5,6 +5,7 @@ use crate::providers::base::{Provider, MSG_COUNT_FOR_SESSION_NAME_GENERATION};
use crate::recipe::Recipe;
use crate::session::extension_data::ExtensionData;
use anyhow::Result;
use chrono::{DateTime, Utc};
use etcetera::{choose_app_strategy, AppStrategy};
use rmcp::model::Role;
use serde::{Deserialize, Serialize};
@@ -27,8 +28,8 @@ pub struct Session {
#[schema(value_type = String)]
pub working_dir: PathBuf,
pub description: String,
pub created_at: String,
pub updated_at: String,
pub created_at: DateTime<Utc>,
pub updated_at: DateTime<Utc>,
pub extension_data: ExtensionData,
pub total_tokens: Option<i32>,
pub input_tokens: Option<i32>,
@@ -279,8 +280,8 @@ impl Default for Session {
id: String::new(),
working_dir: std::env::current_dir().unwrap_or_else(|_| PathBuf::from(".")),
description: String::new(),
created_at: String::new(),
updated_at: String::new(),
created_at: Default::default(),
updated_at: Default::default(),
extension_data: ExtensionData::default(),
total_tokens: None,
input_tokens: None,
@@ -510,8 +511,8 @@ impl SessionStorage {
.bind(&session.id)
.bind(&session.description)
.bind(session.working_dir.to_string_lossy().as_ref())
.bind(&session.created_at)
.bind(&session.updated_at)
.bind(session.created_at)
.bind(session.updated_at)
.bind(serde_json::to_string(&session.extension_data)?)
.bind(session.total_tokens)
.bind(session.input_tokens)
+7 -14
View File
@@ -851,25 +851,18 @@ impl TemporalScheduler {
if let Some(session_info) =
all_sessions.iter().find(|s| s.id == session_id)
{
// Parse the updated_at timestamp from the database
if let Ok(modified_dt) = DateTime::parse_from_str(
&session_info.updated_at,
"%Y-%m-%d %H:%M:%S UTC",
) {
let modified_utc = modified_dt.with_timezone(&Utc);
let now = Utc::now();
let time_diff = now.signed_duration_since(modified_utc);
let now = Utc::now();
let time_diff = now.signed_duration_since(session_info.updated_at);
// Increased tolerance to 5 minutes to reduce false positives
if time_diff.num_minutes() < 5 {
has_active_session = true;
tracing::debug!(
// Increased tolerance to 5 minutes to reduce false positives
if time_diff.num_minutes() < 5 {
has_active_session = true;
tracing::debug!(
"Found active session for job '{}' modified {} minutes ago",
job.id,
time_diff.num_minutes()
);
break;
}
break;
}
}
}
+2 -2
View File
@@ -382,7 +382,7 @@ pub fn create_test_session_metadata(message_count: usize, working_dir: &str) ->
id: "".to_string(),
working_dir: PathBuf::from(working_dir),
description: "Test session".to_string(),
created_at: "".to_string(),
created_at: Default::default(),
schedule_id: Some("test_job".to_string()),
recipe: None,
total_tokens: Some(100),
@@ -392,7 +392,7 @@ pub fn create_test_session_metadata(message_count: usize, working_dir: &str) ->
accumulated_input_tokens: Some(50),
accumulated_output_tokens: Some(50),
extension_data: Default::default(),
updated_at: "".to_string(),
updated_at: Default::default(),
conversation: None,
message_count,
}
+1
View File
@@ -0,0 +1 @@
Temporary file for batch operations
+4 -2
View File
@@ -3650,7 +3650,8 @@
"nullable": true
},
"created_at": {
"type": "string"
"type": "string",
"format": "date-time"
},
"description": {
"type": "string"
@@ -3693,7 +3694,8 @@
"nullable": true
},
"updated_at": {
"type": "string"
"type": "string",
"format": "date-time"
},
"working_dir": {
"type": "string"
@@ -185,6 +185,8 @@ const SessionListView: React.FC<SessionListViewProps> = React.memo(
currentIndex: number;
} | null>(null);
const [visibleGroupsCount, setVisibleGroupsCount] = useState(15);
// Edit modal state
const [showEditModal, setShowEditModal] = useState(false);
const [editingSession, setEditingSession] = useState<Session | null>(null);
@@ -210,6 +212,33 @@ const SessionListView: React.FC<SessionListViewProps> = React.memo(
}
};
const visibleDateGroups = useMemo(() => {
return dateGroups.slice(0, visibleGroupsCount);
}, [dateGroups, visibleGroupsCount]);
const handleScroll = useCallback(
(target: HTMLDivElement) => {
const { scrollTop, scrollHeight, clientHeight } = target;
const threshold = 200;
if (
scrollHeight - scrollTop - clientHeight < threshold &&
visibleGroupsCount < dateGroups.length
) {
setVisibleGroupsCount((prev) => Math.min(prev + 5, dateGroups.length));
}
},
[visibleGroupsCount, dateGroups.length]
);
useEffect(() => {
if (debouncedSearchTerm) {
setVisibleGroupsCount(dateGroups.length);
} else {
setVisibleGroupsCount(15);
}
}, [debouncedSearchTerm, dateGroups.length]);
const loadSessions = useCallback(async () => {
setIsLoading(true);
setShowSkeleton(true);
@@ -557,10 +586,9 @@ const SessionListView: React.FC<SessionListViewProps> = React.memo(
);
}
// For regular rendering in grid layout
return (
<div className="space-y-8">
{dateGroups.map((group) => (
{visibleDateGroups.map((group) => (
<div key={group.label} className="space-y-4">
<div className="sticky top-0 z-10 bg-background-default/95 backdrop-blur-sm">
<h2 className="text-text-muted">{group.label}</h2>
@@ -577,6 +605,15 @@ const SessionListView: React.FC<SessionListViewProps> = React.memo(
</div>
</div>
))}
{visibleGroupsCount < dateGroups.length && (
<div className="flex justify-center py-8">
<div className="flex items-center space-x-2 text-text-muted">
<div className="animate-spin rounded-full h-4 w-4 border-b-2 border-text-muted"></div>
<span>Loading more sessions...</span>
</div>
</div>
)}
</div>
);
};
@@ -597,7 +634,7 @@ const SessionListView: React.FC<SessionListViewProps> = React.memo(
</div>
<div className="flex-1 min-h-0 relative px-8">
<ScrollArea className="h-full" data-search-scroll-area>
<ScrollArea handleScroll={handleScroll} className="h-full" data-search-scroll-area>
<div ref={containerRef} className="h-full relative">
<SearchView
onSearch={handleSearch}
+18 -2
View File
@@ -15,10 +15,22 @@ interface ScrollAreaProps extends React.ComponentPropsWithoutRef<typeof ScrollAr
/* padding needs to be passed into the container inside ScrollArea to avoid pushing the scrollbar out */
paddingX?: number;
paddingY?: number;
handleScroll?: (viewport: HTMLDivElement) => void;
}
const ScrollArea = React.forwardRef<ScrollAreaHandle, ScrollAreaProps>(
({ className, children, autoScroll = false, paddingX, paddingY, ...props }, ref) => {
(
{
className,
children,
autoScroll = false,
paddingX,
paddingY,
handleScroll: handleScrollProp,
...props
},
ref
) => {
const rootRef = React.useRef<React.ElementRef<typeof ScrollAreaPrimitive.Root>>(null);
const viewportRef = React.useRef<HTMLDivElement>(null);
const viewportEndRef = React.useRef<HTMLDivElement>(null);
@@ -71,7 +83,11 @@ const ScrollArea = React.forwardRef<ScrollAreaHandle, ScrollAreaProps>(
setIsFollowing(isAtBottom);
setIsScrolled(scrollTop > 0);
}, []);
if (handleScrollProp) {
handleScrollProp(viewport);
}
}, [handleScrollProp]);
// Track previous scroll height to detect content changes
const prevScrollHeightRef = React.useRef<number>(0);