From d289951b2466e158a1ae01c39c662f1fb60e2d95 Mon Sep 17 00:00:00 2001 From: netdex Date: Sat, 9 Aug 2025 09:25:26 -0400 Subject: [PATCH] GPT-5 support --- niinii/src/app.rs | 13 ++--- niinii/src/settings.rs | 2 + niinii/src/translator/chat.rs | 69 ++++++++++++---------- niinii/src/view/translator/chat.rs | 21 +++++-- openai-chat/src/chat/chat_buffer.rs | 6 +- openai-chat/src/chat/mod.rs | 91 +++++++++++++++++++---------- openai-chat/src/protocol/chat.rs | 28 +++++++-- openai-chat/src/realtime/mod.rs | 8 +-- 8 files changed, 153 insertions(+), 85 deletions(-) diff --git a/niinii/src/app.rs b/niinii/src/app.rs index 4dee34e..06131a3 100644 --- a/niinii/src/app.rs +++ b/niinii/src/app.rs @@ -165,7 +165,7 @@ impl App { } } - fn transition(&mut self, ui: &Ui, state: State) { + fn transition(&mut self, _ui: &Ui, state: State) { self.state = state; } @@ -213,14 +213,11 @@ impl App { } } - match &self.state { - State::Completed => { - if let Some(request_gloss_text) = self.request_gloss_text.clone() { - self.request_gloss_text = None; - self.request_parse(ui, &request_gloss_text); - } + if let State::Completed = &self.state { + if let Some(request_gloss_text) = self.request_gloss_text.clone() { + self.request_gloss_text = None; + self.request_parse(ui, &request_gloss_text); } - _ => (), }; if self.settings.watch_clipboard { diff --git a/niinii/src/settings.rs b/niinii/src/settings.rs index 2cd8bfd..ea3868c 100644 --- a/niinii/src/settings.rs +++ b/niinii/src/settings.rs @@ -36,6 +36,7 @@ pub struct ChatSettings { pub presence_penalty: Option, pub connection_timeout: u64, pub timeout: u64, + pub stream: bool, } impl Default for ChatSettings { fn default() -> Self { @@ -51,6 +52,7 @@ impl Default for ChatSettings { presence_penalty: None, connection_timeout: 3000, timeout: 10000, + stream: true, } } } diff --git a/niinii/src/translator/chat.rs b/niinii/src/translator/chat.rs index 26b109c..8301253 100644 --- a/niinii/src/translator/chat.rs +++ b/niinii/src/translator/chat.rs @@ -53,7 +53,7 @@ impl Translator for ChatTranslator { let chatgpt = &settings.chat; let permit = self.semaphore.clone().acquire_owned().await.unwrap(); - let mut exchange = { + let exchange = { let mut chat = self.chat.lock().await; chat.start_exchange( Message { @@ -74,47 +74,54 @@ impl Translator for ChatTranslator { messages: exchange.prompt(), temperature: chatgpt.temperature, top_p: chatgpt.top_p, - max_tokens: chatgpt.max_tokens, + max_completion_tokens: chatgpt.max_tokens, presence_penalty: chatgpt.presence_penalty, ..Default::default() }; let exchange = Arc::new(Mutex::new(exchange)); - let mut stream = self.client.stream(chat_request).await?; let token = CancellationToken::new(); - let chat = &self.chat; - tokio::spawn( - enclose! { (chat, token, exchange, chatgpt.max_context_tokens => max_context_tokens) async move { - // Hold permit: We are not allowed to begin another translation - // request until this one is complete. - let _permit = permit; - loop { - tokio::select! { - msg = stream.next() => match msg { - Some(Ok(completion)) => { - let mut exchange = exchange.lock().await; - let message = &completion.choices.first().unwrap().delta; - exchange.append(message) + if chatgpt.stream { + let mut stream = self.client.stream(chat_request).await?; + let chat = &self.chat; + tokio::spawn( + enclose! { (chat, token, exchange, chatgpt.max_context_tokens => max_context_tokens) async move { + // Hold permit: We are not allowed to begin another translation + // request until this one is complete. + let _permit = permit; + loop { + tokio::select! { + msg = stream.next() => match msg { + Some(Ok(completion)) => { + let mut exchange = exchange.lock().await; + let message = &completion.choices.first().unwrap().delta; + exchange.append_partial(message) + }, + Some(Err(err)) => { + tracing::error!(%err, "stream"); + break + }, + None => { + let mut chat = chat.lock().await; + let mut exchange = exchange.lock().await; + chat.commit(&mut exchange); + chat.enforce_context_limit(max_context_tokens); + break + } }, - Some(Err(err)) => { - tracing::error!(%err, "stream"); - break - }, - None => { - let mut chat = chat.lock().await; - let mut exchange = exchange.lock().await; - chat.commit(&mut exchange); - chat.enforce_context_limit(max_context_tokens); + _ = token.cancelled() => { break } - }, - _ = token.cancelled() => { - break } } - } - }.instrument(tracing::Span::current())}, - ); + }.instrument(tracing::Span::current())}, + ); + } else { + let response = self.client.chat(chat_request).await?; + let mut e = exchange.lock().await; + e.set_complete(response.choices.first().unwrap().message.clone()); + self.chat.lock().await.commit(&mut e); + } Ok(Box::new(ChatTranslation { model: chatgpt.model, diff --git a/niinii/src/view/translator/chat.rs b/niinii/src/view/translator/chat.rs index 09385ef..4cd8574 100644 --- a/niinii/src/view/translator/chat.rs +++ b/niinii/src/view/translator/chat.rs @@ -4,7 +4,7 @@ use openai_chat::chat::{Model, Role}; use crate::{ settings::Settings, translator::chat::{ChatTranslation, ChatTranslator}, - view::mixins::drag_handle, + view::mixins::{drag_handle, help_marker}, }; use crate::view::{ @@ -30,7 +30,10 @@ impl View for ViewChatTranslator<'_> { } }); if ui.collapsing_header("Tuning", TreeNodeFlags::DEFAULT_OPEN) { - if let Some(_token) = ui.begin_table("##", 2) { + if let Some(_token) = ui.begin_table("##", 3) { + ui.table_next_column(); + ui.set_next_item_width(ui.current_font_size() * -8.0); + combo_enum(ui, "Model", &mut chatgpt.model); ui.table_next_column(); ui.set_next_item_width(ui.current_font_size() * -8.0); ui.input_scalar("Max context tokens", &mut chatgpt.max_context_tokens) @@ -75,8 +78,12 @@ impl View for ViewChatTranslator<'_> { }, ); ui.table_next_column(); - ui.set_next_item_width(ui.current_font_size() * -8.0); - combo_enum(ui, "Model", &mut chatgpt.model); + ui.checkbox("Stream", &mut chatgpt.stream); + ui.same_line(); + help_marker( + ui, + "Use streaming API (may require ID verification for some models)", + ); } } ui.child_window("context_window").build(|| { @@ -194,13 +201,15 @@ impl View for ViewChatTranslation<'_> { let _wrap_token = ui.push_text_wrap_pos_with_pos(0.0); ui.text(""); // anchor for line wrapping ui.same_line(); - let ChatTranslation { exchange, .. } = self.0; + let ChatTranslation { + model, exchange, .. + } = self.0; let exchange = exchange.blocking_lock(); let draw_list = ui.get_window_draw_list(); stroke_text_with_highlight( ui, &draw_list, - "[ChatGPT]", + &format!("[{}]", std::convert::Into::<&'static str>::into(model)), 1.0, Some(StyleColor::NavHighlight), ); diff --git a/openai-chat/src/chat/chat_buffer.rs b/openai-chat/src/chat/chat_buffer.rs index 19ee98f..4d9386e 100644 --- a/openai-chat/src/chat/chat_buffer.rs +++ b/openai-chat/src/chat/chat_buffer.rs @@ -89,7 +89,7 @@ pub struct Exchange { usage: Option, } impl Exchange { - pub fn append(&mut self, partial: &PartialMessage) { + pub fn append_partial(&mut self, partial: &PartialMessage) { if let Some(last) = &mut self.response { if let Some(content) = &mut last.content { content.push_str(&partial.content) @@ -104,6 +104,10 @@ impl Exchange { } } + pub fn set_complete(&mut self, message: Message) { + self.response = Some(message); + } + pub fn prompt(&self) -> Vec { let mut messages = vec![]; messages.push(self.system.clone()); diff --git a/openai-chat/src/chat/mod.rs b/openai-chat/src/chat/mod.rs index 82391a1..6591f3b 100644 --- a/openai-chat/src/chat/mod.rs +++ b/openai-chat/src/chat/mod.rs @@ -11,7 +11,7 @@ pub use crate::protocol::chat::{Message, Model, PartialMessage, Role, Usage}; pub use chat_buffer::{ChatBuffer, Exchange}; use crate::{ - protocol::chat::{self, StreamResponse}, + protocol::chat::{self, ChatResponse, StreamOptions, StreamResponse}, Client, Error, }; @@ -37,7 +37,7 @@ pub struct Request { /// The maximum number of tokens allowed for the generated answer. By /// default, the number of tokens the model can return will be (4096 - prompt /// tokens). - pub max_tokens: Option, + pub max_completion_tokens: Option, /// Number between -2.0 and 2.0. Positive values penalize new tokens based /// on whether they appear in the text so far, increasing the model's /// likelihood to talk about new topics. @@ -51,7 +51,7 @@ impl From for chat::Request { temperature, top_p, n, - max_tokens, + max_completion_tokens, presence_penalty, } = value; chat::Request { @@ -60,7 +60,7 @@ impl From for chat::Request { temperature, top_p, n, - max_tokens, + max_completion_tokens, presence_penalty, ..Default::default() } @@ -71,7 +71,7 @@ impl Client { #[tracing::instrument(level = Level::DEBUG, skip_all, err)] pub async fn chat(&self, request: Request) -> Result { let request: chat::Request = request.into(); - tracing::trace!(?request); + tracing::debug!(?request); let response: chat::ChatResponse = self .shared .request(Method::POST, "/v1/chat/completions") @@ -80,7 +80,7 @@ impl Client { .await? .json() .await?; - tracing::trace!(?response); + tracing::debug!(?response); Ok(response.0?) } @@ -91,39 +91,68 @@ impl Client { ) -> Result>, Error> { let mut request: chat::Request = request.into(); request.stream = Some(true); + request.stream_options = Some(StreamOptions { + include_obfuscation: false, + include_usage: false, + }); - tracing::trace!(?request); - let stream = self + tracing::debug!(?request); + let response = self .shared .request(Method::POST, "/v1/chat/completions") .body(&request) .send() - .await? - .bytes_stream() - .eventsource(); - Ok(stream.map_while(|event| match event { - Ok(event) => { - if event.data == "[DONE]" { - None - } else { - let response = match serde_json::from_str::(&event.data) { - Ok(response) => { - tracing::trace!(?response); - Ok::<_, Error>(response.0) - } - Err(err) => { - tracing::error!(?err, ?event.data); - Err(err.into()) + .await?; + let status = response.status(); + + if status.is_success() { + // HTTP success: Expect SSE response + let stream = response.bytes_stream().eventsource(); + Ok(stream.map_while(|event| { + tracing::trace!(?event); + match event { + Ok(event) => { + if event.data == "[DONE]" { + None + } else { + let response = match serde_json::from_str::(&event.data) + { + Ok(response) => { + tracing::debug!(?response); + Ok::<_, Error>(response.0) + } + Err(err) => { + // Serde error + tracing::error!(?err, ?event.data); + Err(err.into()) + } + }; + Some(response) } - }; - Some(response) + } + Err(err) => { + // SSE error + tracing::error!(?err); + Some(Err(err.into())) + } + } + })) + } else { + // HTTP error: Expect JSON response + let response_err = response.error_for_status_ref().unwrap_err(); + let chat_response = response.json::().await; + match chat_response { + Ok(err) => { + // OpenAI application error + Err(Error::Protocol(err.0.unwrap_err())) + } + Err(err) => { + // Not application error, return HTTP error + tracing::error!(?response_err, ?err, "unexpected stream response"); + Err(response_err.into()) } } - Err(err) => { - tracing::error!(?err); - Some(Err(err.into())) - } - })) + } } } diff --git a/openai-chat/src/protocol/chat.rs b/openai-chat/src/protocol/chat.rs index ec514b1..f6fa98f 100644 --- a/openai-chat/src/protocol/chat.rs +++ b/openai-chat/src/protocol/chat.rs @@ -28,7 +28,6 @@ pub enum Model { #[serde(rename = "gpt-4.1-2025-04-14")] Gpt41_20250414, - #[default] #[serde(rename = "gpt-4.1-mini")] Gpt41Mini, #[serde(rename = "gpt-4.1-mini-2025-04-14")] @@ -38,6 +37,22 @@ pub enum Model { Gpt41Nano, #[serde(rename = "gpt-4.1-nano-2025-04-14")] Gpt41Nano20250414, + + #[serde(rename = "gpt-5")] + Gpt5, + #[serde(rename = "gpt-5-2025-08-07")] + Gpt5_20250807, + + #[default] + #[serde(rename = "gpt-5-mini")] + Gpt5Mini, + #[serde(rename = "gpt-5-mini-2025-08-07")] + Gpt5Mini20250807, + + #[serde(rename = "gpt-5-nano")] + Gpt5Nano, + #[serde(rename = "gpt-5-nano-2025-08-07")] + Gpt5Nano20250807, } #[derive(Debug, Clone, Default, Deserialize, Serialize, PartialEq, Eq, IntoStaticStr, EnumIter)] @@ -47,7 +62,6 @@ pub enum Role { #[default] User, Assistant, - Function, } #[serde_with::skip_serializing_none] @@ -81,6 +95,12 @@ impl Default for Message { } } +#[derive(Debug, Clone, Serialize)] +pub struct StreamOptions { + pub include_obfuscation: bool, + pub include_usage: bool, +} + #[serde_with::skip_serializing_none] #[derive(Debug, Clone, Default, Serialize)] pub struct Request { @@ -90,12 +110,12 @@ pub struct Request { pub top_p: Option, pub n: Option, pub stream: Option, + pub stream_options: Option, #[serde(skip_serializing_if = "Vec::is_empty")] pub stop: Vec, - pub max_tokens: Option, + pub max_completion_tokens: Option, pub presence_penalty: Option, // logit_bias - // user } #[derive(Debug, Clone, Deserialize)] diff --git a/openai-chat/src/realtime/mod.rs b/openai-chat/src/realtime/mod.rs index cd9de93..d584c91 100644 --- a/openai-chat/src/realtime/mod.rs +++ b/openai-chat/src/realtime/mod.rs @@ -299,7 +299,7 @@ impl Client { session_parameters: SessionParameters, ) -> Result { let create_session_request = CreateSessionRequest(session_parameters); - tracing::trace!(?create_session_request); + tracing::debug!(?create_session_request); let create_session_response: CreateSessionResponse = self .shared .request(Method::POST, "/v1/realtime/sessions") @@ -309,7 +309,7 @@ impl Client { .await? .json() .await?; - tracing::trace!(?create_session_response); + tracing::debug!(?create_session_response); RealtimeSession::new(create_session_response).await } } @@ -327,7 +327,7 @@ where if let tungstenite::Message::Text(_) = message { let response = (&message) .try_into() - .inspect(|response| tracing::trace!(?response)) + .inspect(|response| tracing::debug!(?response)) .inspect_err(|err| tracing::error!(%err)); return Ok(Some(response?)); } else { @@ -346,7 +346,7 @@ where S: Sink + Unpin, { async fn send_request(&mut self, request: &ClientEventRequest) -> Result<(), Error> { - tracing::trace!(?request); + tracing::debug!(?request); Ok(self.send(request.try_into()?).await?) } }