feat: better chat support for thinking

This commit is contained in:
Georg Bauer
2026-07-24 18:01:53 +02:00
parent ff984bb96a
commit 5ab110ed0b
9 changed files with 463 additions and 68 deletions

View File

@@ -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");
}
}