mirror of
https://github.com/aaif-goose/goose.git
synced 2026-07-03 14:10:03 +02:00
Merge branch 'main' into platform-extensions
This commit is contained in:
@@ -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),
|
||||
|
||||
@@ -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.
|
||||
}
|
||||
|
||||
@@ -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");
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
|
||||
@@ -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;
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
|
||||
@@ -0,0 +1 @@
|
||||
Temporary file for batch operations
|
||||
@@ -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}
|
||||
|
||||
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user