feat: better chat support for thinking
This commit is contained in:
168
src/app.rs
168
src/app.rs
@@ -2,7 +2,7 @@ mod view;
|
||||
|
||||
pub(crate) use view::app_theme;
|
||||
|
||||
use crate::database::{AppPreferences, Database, ProjectWithSessions};
|
||||
use crate::database::{AppPreferences, Database, ProjectWithSessions, StoredMessage};
|
||||
#[cfg(target_os = "macos")]
|
||||
use crate::engine::{ChatTurn, Generator};
|
||||
use crate::model::{self, DownloadOutcome, DownloadProgress, ManagedArtifactId, ModelChoice};
|
||||
@@ -410,11 +410,38 @@ pub(crate) struct App {
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub(super) struct ChatMessage {
|
||||
id: i32,
|
||||
pub(super) user: bool,
|
||||
reasoning: bool,
|
||||
pub(super) reasoning: Option<String>,
|
||||
reasoning_complete: bool,
|
||||
pub(super) reasoning_open: bool,
|
||||
pub(super) content: String,
|
||||
}
|
||||
|
||||
impl ChatMessage {
|
||||
fn append(&mut self, reasoning: bool, chunk: &str) {
|
||||
if reasoning {
|
||||
self.reasoning.get_or_insert_default().push_str(chunk);
|
||||
} else {
|
||||
self.reasoning_complete = true;
|
||||
self.content.push_str(chunk);
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<StoredMessage> for ChatMessage {
|
||||
fn from(message: StoredMessage) -> Self {
|
||||
Self {
|
||||
id: message.id,
|
||||
user: message.user,
|
||||
reasoning: message.reasoning,
|
||||
reasoning_complete: message.reasoning_complete,
|
||||
reasoning_open: false,
|
||||
content: message.content,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(target_os = "macos")]
|
||||
struct GenerationWorker {
|
||||
commands: mpsc::Sender<GenerationCommand>,
|
||||
@@ -436,7 +463,7 @@ enum GenerationCommand {
|
||||
#[cfg(target_os = "macos")]
|
||||
enum GenerationEvent {
|
||||
Loading,
|
||||
Chunk(String),
|
||||
Chunk { reasoning: bool, content: String },
|
||||
Finished(Result<(), String>),
|
||||
}
|
||||
|
||||
@@ -519,6 +546,7 @@ pub(crate) enum Message {
|
||||
StopModelDownload,
|
||||
DownloadProgressTick,
|
||||
ComposerChanged(String),
|
||||
ToggleReasoning(usize),
|
||||
SubmitPrompt,
|
||||
StopGeneration,
|
||||
GenerationTick,
|
||||
@@ -870,6 +898,13 @@ impl App {
|
||||
}
|
||||
Message::DownloadProgressTick => self.update_download_progress(),
|
||||
Message::ComposerChanged(value) => self.composer = value,
|
||||
Message::ToggleReasoning(index) => {
|
||||
if let Some(message) = self.conversation.get_mut(index)
|
||||
&& message.reasoning.is_some()
|
||||
{
|
||||
message.reasoning_open = !message.reasoning_open;
|
||||
}
|
||||
}
|
||||
Message::SubmitPrompt => self.start_generation(),
|
||||
Message::StopGeneration => {
|
||||
#[cfg(target_os = "macos")]
|
||||
@@ -916,6 +951,8 @@ impl App {
|
||||
}
|
||||
self.selected_project = Some(project_id);
|
||||
self.selected_session = None;
|
||||
self.conversation.clear();
|
||||
self.composer.clear();
|
||||
self.error = None;
|
||||
}
|
||||
Message::DeleteProject(project_id) => {
|
||||
@@ -930,6 +967,8 @@ impl App {
|
||||
if self.selected_project == Some(project_id) {
|
||||
self.selected_project = None;
|
||||
self.selected_session = None;
|
||||
self.conversation.clear();
|
||||
self.composer.clear();
|
||||
}
|
||||
self.reload_projects();
|
||||
}
|
||||
@@ -944,13 +983,24 @@ impl App {
|
||||
Some("Stop the active generation before changing sessions.".into());
|
||||
return Task::none();
|
||||
}
|
||||
if self.selected_session != Some(session_id) {
|
||||
self.conversation.clear();
|
||||
self.composer.clear();
|
||||
if self.selected_session == Some(session_id) {
|
||||
return Task::none();
|
||||
}
|
||||
let Some(database) = &mut self.database else {
|
||||
return Task::none();
|
||||
};
|
||||
match database.load_messages(session_id) {
|
||||
Ok(messages) => {
|
||||
self.conversation = messages.into_iter().map(ChatMessage::from).collect();
|
||||
self.composer.clear();
|
||||
self.selected_project = Some(project_id);
|
||||
self.selected_session = Some(session_id);
|
||||
self.error = None;
|
||||
}
|
||||
Err(error) => {
|
||||
self.error = Some(format!("Could not load the chat session: {error}"));
|
||||
}
|
||||
}
|
||||
self.selected_project = Some(project_id);
|
||||
self.selected_session = Some(session_id);
|
||||
self.error = None;
|
||||
}
|
||||
Message::DeleteSession(session_id) => {
|
||||
if self.generating && self.selected_session == Some(session_id) {
|
||||
@@ -963,6 +1013,8 @@ impl App {
|
||||
Ok(()) => {
|
||||
if self.selected_session == Some(session_id) {
|
||||
self.selected_session = None;
|
||||
self.conversation.clear();
|
||||
self.composer.clear();
|
||||
}
|
||||
self.reload_projects();
|
||||
}
|
||||
@@ -1223,6 +1275,8 @@ impl App {
|
||||
match database.create_session(project_id, &title) {
|
||||
Ok(session) => {
|
||||
self.selected_session = Some(session.id);
|
||||
self.conversation.clear();
|
||||
self.composer.clear();
|
||||
self.error = None;
|
||||
self.reload_projects();
|
||||
}
|
||||
@@ -1333,20 +1387,25 @@ impl App {
|
||||
}
|
||||
};
|
||||
let assistant_reasoning = effective.turn.reasoning_mode != ReasoningMode::Direct;
|
||||
let session_id = self
|
||||
.selected_session
|
||||
.expect("a selected session was checked");
|
||||
#[cfg(target_os = "macos")]
|
||||
let mut messages = self
|
||||
.conversation
|
||||
.iter()
|
||||
.map(|message| ChatTurn {
|
||||
user: message.user,
|
||||
reasoning: message.reasoning,
|
||||
reasoning: message.reasoning.clone(),
|
||||
reasoning_complete: message.reasoning_complete,
|
||||
content: message.content.clone(),
|
||||
})
|
||||
.collect::<Vec<_>>();
|
||||
#[cfg(target_os = "macos")]
|
||||
messages.push(ChatTurn {
|
||||
user: true,
|
||||
reasoning: false,
|
||||
reasoning: None,
|
||||
reasoning_complete: true,
|
||||
content: prompt.clone(),
|
||||
});
|
||||
|
||||
@@ -1361,6 +1420,16 @@ impl App {
|
||||
}
|
||||
}
|
||||
}
|
||||
let Some(database) = &mut self.database else {
|
||||
return;
|
||||
};
|
||||
let saved = match database.start_chat_turn(session_id, &prompt, assistant_reasoning) {
|
||||
Ok(turn) => turn,
|
||||
Err(error) => {
|
||||
self.error = Some(format!("Could not save the chat turn: {error}"));
|
||||
return;
|
||||
}
|
||||
};
|
||||
let cancel = Arc::new(AtomicBool::new(false));
|
||||
let idle_timeout =
|
||||
Duration::from_secs(self.preferences.idle_timeout_minutes.max(1) as u64 * 60);
|
||||
@@ -1378,6 +1447,14 @@ impl App {
|
||||
return;
|
||||
}
|
||||
worker.cancel = Some(cancel);
|
||||
let user = ChatMessage::from(saved.0);
|
||||
let mut assistant = ChatMessage::from(saved.1);
|
||||
assistant.reasoning_open = assistant_reasoning;
|
||||
self.composer.clear();
|
||||
self.conversation.push(user);
|
||||
self.conversation.push(assistant);
|
||||
self.generating = true;
|
||||
self.error = None;
|
||||
}
|
||||
#[cfg(not(target_os = "macos"))]
|
||||
{
|
||||
@@ -1385,19 +1462,6 @@ impl App {
|
||||
self.error = Some("Local Metal generation requires macOS.".into());
|
||||
return;
|
||||
}
|
||||
self.composer.clear();
|
||||
self.conversation.push(ChatMessage {
|
||||
user: true,
|
||||
reasoning: false,
|
||||
content: prompt,
|
||||
});
|
||||
self.conversation.push(ChatMessage {
|
||||
user: false,
|
||||
reasoning: assistant_reasoning,
|
||||
content: String::new(),
|
||||
});
|
||||
self.generating = true;
|
||||
self.error = None;
|
||||
}
|
||||
|
||||
fn poll_generation(&mut self) {
|
||||
@@ -1407,14 +1471,17 @@ impl App {
|
||||
return;
|
||||
};
|
||||
#[cfg(target_os = "macos")]
|
||||
let mut transcript_changed = false;
|
||||
#[cfg(target_os = "macos")]
|
||||
loop {
|
||||
match worker.events.try_recv() {
|
||||
Ok(GenerationEvent::Loading) => {}
|
||||
Ok(GenerationEvent::Chunk(chunk)) => {
|
||||
Ok(GenerationEvent::Chunk { reasoning, content }) => {
|
||||
if let Some(message) = self.conversation.last_mut()
|
||||
&& !message.user
|
||||
{
|
||||
message.content.push_str(&chunk);
|
||||
message.append(reasoning, &content);
|
||||
transcript_changed = true;
|
||||
}
|
||||
}
|
||||
Ok(GenerationEvent::Finished(result)) => {
|
||||
@@ -1434,6 +1501,26 @@ impl App {
|
||||
}
|
||||
}
|
||||
}
|
||||
#[cfg(target_os = "macos")]
|
||||
if transcript_changed
|
||||
&& let Some(message) = self.conversation.last()
|
||||
&& let Some(database) = &mut self.database
|
||||
&& let Err(error) = database.update_message(
|
||||
message.id,
|
||||
message.reasoning.as_deref(),
|
||||
message.reasoning_complete,
|
||||
&message.content,
|
||||
)
|
||||
{
|
||||
if let Some(cancel) = self
|
||||
.generation_worker
|
||||
.as_ref()
|
||||
.and_then(|worker| worker.cancel.as_ref())
|
||||
{
|
||||
cancel.store(true, Ordering::Relaxed);
|
||||
}
|
||||
self.error = Some(format!("Could not save generated chat text: {error}"));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1472,9 +1559,15 @@ fn spawn_generation_worker() -> Result<GenerationWorker, String> {
|
||||
};
|
||||
}
|
||||
if let Some((_, generator)) = &mut loaded {
|
||||
let result = generator.generate(&messages, &turn, &cancel, |chunk| {
|
||||
let _ = event_sender.send(GenerationEvent::Chunk(chunk));
|
||||
});
|
||||
let result = generator.generate(
|
||||
&messages,
|
||||
&turn,
|
||||
&cancel,
|
||||
|reasoning, content| {
|
||||
let _ = event_sender
|
||||
.send(GenerationEvent::Chunk { reasoning, content });
|
||||
},
|
||||
);
|
||||
let _ = event_sender.send(GenerationEvent::Finished(result));
|
||||
last_used = Instant::now();
|
||||
}
|
||||
@@ -1588,4 +1681,21 @@ mod tests {
|
||||
assert!(parse_streaming_cache("1.5GB").is_err());
|
||||
assert_eq!(parse_optional_gib("Memory", "8GB").unwrap(), Some(8));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn assistant_stream_keeps_reasoning_separate_from_the_answer() {
|
||||
let mut message = ChatMessage {
|
||||
id: 1,
|
||||
user: false,
|
||||
reasoning: Some(String::new()),
|
||||
reasoning_complete: false,
|
||||
reasoning_open: true,
|
||||
content: String::new(),
|
||||
};
|
||||
message.append(true, "working it out");
|
||||
message.append(false, "final answer");
|
||||
assert_eq!(message.reasoning.as_deref(), Some("working it out"));
|
||||
assert!(message.reasoning_complete);
|
||||
assert_eq!(message.content, "final answer");
|
||||
}
|
||||
}
|
||||
|
||||
Reference in New Issue
Block a user