diff options
| author | sigoden <sigoden@gmail.com> | 2024-10-19 18:25:40 +0800 |
|---|---|---|
| committer | GitHub <noreply@github.com> | 2024-10-19 18:25:40 +0800 |
| commit | 7e29f642a99682ea7bb074eb4e7aa7d34d523515 (patch) | |
| tree | 248bbcb5d3af83981e451eb115f1fde23c8386de /src | |
| parent | c5e4421e0da6d647a636ec4cf23c110ff515e79a (diff) | |
| download | aichat-7e29f642a99682ea7bb074eb4e7aa7d34d523515.tar.gz | |
feat: support openai o1 models (#935)
Diffstat (limited to 'src')
| -rw-r--r-- | src/client/model.rs | 12 | ||||
| -rw-r--r-- | src/client/openai.rs | 6 | ||||
| -rw-r--r-- | src/config/input.rs | 4 | ||||
| -rw-r--r-- | src/main.rs | 4 | ||||
| -rw-r--r-- | src/repl/mod.rs | 2 |
5 files changed, 24 insertions, 4 deletions
diff --git a/src/client/model.rs b/src/client/model.rs index 08ced10..f8a23ff 100644 --- a/src/client/model.rs +++ b/src/client/model.rs @@ -183,6 +183,14 @@ impl Model { self.data.supports_vision } + pub fn no_stream(&self) -> bool { + self.data.no_stream + } + + pub fn no_system_message(&self) -> bool { + self.data.no_system_message + } + pub fn max_tokens_per_chunk(&self) -> Option<usize> { self.data.max_tokens_per_chunk } @@ -268,6 +276,10 @@ pub struct ModelData { pub supports_vision: bool, #[serde(default)] pub supports_function_calling: bool, + #[serde(default)] + no_stream: bool, + #[serde(default)] + no_system_message: bool, // embedding-only properties pub max_tokens_per_chunk: Option<usize>, diff --git a/src/client/openai.rs b/src/client/openai.rs index c4c2b0c..831b64f 100644 --- a/src/client/openai.rs +++ b/src/client/openai.rs @@ -193,13 +193,17 @@ struct EmbeddingsResBodyEmbedding { pub fn openai_build_chat_completions_body(data: ChatCompletionsData, model: &Model) -> Value { let ChatCompletionsData { - messages, + mut messages, temperature, top_p, functions, stream, } = data; + if model.no_system_message() { + patch_system_message(&mut messages); + } + let messages: Vec<Value> = messages .into_iter() .flat_map(|message| { diff --git a/src/config/input.rs b/src/config/input.rs index 7662d2c..0325e1c 100644 --- a/src/config/input.rs +++ b/src/config/input.rs @@ -135,6 +135,10 @@ impl Input { self.text = text; } + pub fn stream(&self) -> bool { + self.config.read().stream && !self.role().model().no_stream() + } + pub fn continue_output(&self) -> Option<&str> { self.continue_output.as_deref() } diff --git a/src/main.rs b/src/main.rs index ec61370..a7cc9f1 100644 --- a/src/main.rs +++ b/src/main.rs @@ -174,7 +174,7 @@ async fn start_directive( let client = input.create_client()?; let extract_code = !*IS_STDOUT_TERMINAL && code_mode; config.write().before_chat_completion(&input)?; - let (output, tool_results) = if !config.read().stream || extract_code { + let (output, tool_results) = if !input.stream() || extract_code { let task = client.chat_completions(input.clone()); let ret = run_with_spinner(task, "Generating").await; match ret { @@ -288,7 +288,7 @@ async fn shell_execute(config: &GlobalConfig, shell: &Shell, mut input: Input) - let role = config.read().retrieve_role(EXPLAIN_SHELL_ROLE)?; let input = Input::from_str(config, &eval_str, Some(role)); let abort = create_abort_signal(); - if config.read().stream { + if input.stream() { call_chat_completions_streaming(&input, client.as_ref(), config, abort) .await?; } else { diff --git a/src/repl/mod.rs b/src/repl/mod.rs index d1068fd..735ad25 100644 --- a/src/repl/mod.rs +++ b/src/repl/mod.rs @@ -657,7 +657,7 @@ async fn ask( let client = input.create_client()?; config.write().before_chat_completion(&input)?; - let (output, tool_results) = if config.read().stream { + let (output, tool_results) = if input.stream() { call_chat_completions_streaming(&input, client.as_ref(), config, abort_signal.clone()) .await? } else { |
