Implement OpenAI's response API and add support for built-in tools (web search & code interpreter).
This commit is contained in:
@@ -17,7 +17,7 @@ path = "src/lib.rs"
|
|||||||
[dependencies]
|
[dependencies]
|
||||||
anthropic = { git = "https://github.com/etkecc/anthropic-rs.git", branch = "fix-content-block-image" }
|
anthropic = { git = "https://github.com/etkecc/anthropic-rs.git", branch = "fix-content-block-image" }
|
||||||
anyhow = "1.0.*"
|
anyhow = "1.0.*"
|
||||||
async-openai = { version = "0.32.3", features = ["audio", "chat-completion", "image"] }
|
async-openai = { version = "0.32.3", features = ["audio", "chat-completion", "image", "responses"] }
|
||||||
base64 = "0.22.*"
|
base64 = "0.22.*"
|
||||||
chrono = { version = "0.4.*", default-features = false, features = ["std", "now"] }
|
chrono = { version = "0.4.*", default-features = false, features = ["std", "now"] }
|
||||||
# We'd rather not depend on this, but we cannot use the ruma-events EventContent macro without it.
|
# We'd rather not depend on this, but we cannot use the ruma-events EventContent macro without it.
|
||||||
|
|||||||
@@ -9,6 +9,10 @@ text_generation:
|
|||||||
max_response_tokens: null
|
max_response_tokens: null
|
||||||
max_completion_tokens: 128000
|
max_completion_tokens: 128000
|
||||||
max_context_tokens: 400000
|
max_context_tokens: 400000
|
||||||
|
# Built-in tools
|
||||||
|
tools:
|
||||||
|
web_search: false
|
||||||
|
code_interpreter: false
|
||||||
speech_to_text:
|
speech_to_text:
|
||||||
model_id: whisper-1
|
model_id: whisper-1
|
||||||
text_to_speech:
|
text_to_speech:
|
||||||
|
|||||||
@@ -64,6 +64,9 @@ pub struct TextGenerationConfig {
|
|||||||
|
|
||||||
#[serde(default)]
|
#[serde(default)]
|
||||||
pub max_context_tokens: u32,
|
pub max_context_tokens: u32,
|
||||||
|
|
||||||
|
#[serde(default)]
|
||||||
|
pub tools: ToolsConfig,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Default for TextGenerationConfig {
|
impl Default for TextGenerationConfig {
|
||||||
@@ -75,6 +78,7 @@ impl Default for TextGenerationConfig {
|
|||||||
max_response_tokens: None,
|
max_response_tokens: None,
|
||||||
max_completion_tokens: Some(128_000),
|
max_completion_tokens: Some(128_000),
|
||||||
max_context_tokens: 400_000,
|
max_context_tokens: 400_000,
|
||||||
|
tools: ToolsConfig::default(),
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -83,6 +87,15 @@ fn default_text_model_id() -> String {
|
|||||||
"gpt-5.2".to_owned()
|
"gpt-5.2".to_owned()
|
||||||
}
|
}
|
||||||
|
|
||||||
|
#[derive(Debug, Clone, Serialize, Deserialize, Default)]
|
||||||
|
pub struct ToolsConfig {
|
||||||
|
#[serde(default)]
|
||||||
|
pub web_search: bool,
|
||||||
|
|
||||||
|
#[serde(default)]
|
||||||
|
pub code_interpreter: bool,
|
||||||
|
}
|
||||||
|
|
||||||
#[derive(Debug, Clone, Serialize, Deserialize)]
|
#[derive(Debug, Clone, Serialize, Deserialize)]
|
||||||
pub struct SpeechToTextConfig {
|
pub struct SpeechToTextConfig {
|
||||||
#[serde(default = "default_speech_to_text_model_id")]
|
#[serde(default = "default_speech_to_text_model_id")]
|
||||||
|
|||||||
@@ -5,7 +5,10 @@ use async_openai::{
|
|||||||
config::OpenAIConfig,
|
config::OpenAIConfig,
|
||||||
types::{
|
types::{
|
||||||
audio::{AudioInput, CreateSpeechRequestArgs, CreateTranscriptionRequestArgs},
|
audio::{AudioInput, CreateSpeechRequestArgs, CreateTranscriptionRequestArgs},
|
||||||
chat::{ChatCompletionRequestMessage, CreateChatCompletionRequestArgs},
|
responses::{
|
||||||
|
CodeInterpreterContainerAuto, CodeInterpreterTool, CodeInterpreterToolContainer,
|
||||||
|
CreateResponseArgs, OutputItem, OutputMessageContent, Tool, WebSearchTool,
|
||||||
|
},
|
||||||
images::{
|
images::{
|
||||||
CreateImageEditRequestArgs, CreateImageRequestArgs,
|
CreateImageEditRequestArgs, CreateImageRequestArgs,
|
||||||
Image, ImageInput, ImageModel, ImageResponseFormat,
|
Image, ImageInput, ImageModel, ImageResponseFormat,
|
||||||
@@ -129,28 +132,44 @@ impl ControllerTrait for Controller {
|
|||||||
conversation_messages.insert(0, prompt_message);
|
conversation_messages.insert(0, prompt_message);
|
||||||
}
|
}
|
||||||
|
|
||||||
let openai_conversation_messages: Vec<ChatCompletionRequestMessage> =
|
let input = super::utils::convert_llm_messages_to_openai_response_input(conversation_messages);
|
||||||
super::utils::convert_llm_messages_to_openai_messages(conversation_messages);
|
|
||||||
|
|
||||||
let messages_count = openai_conversation_messages.len();
|
let messages_count = match &input {
|
||||||
|
async_openai::types::responses::InputParam::Items(items) => items.len(),
|
||||||
|
_ => 1,
|
||||||
|
};
|
||||||
|
|
||||||
let temperature = params
|
let temperature = params
|
||||||
.temperature_override
|
.temperature_override
|
||||||
.unwrap_or(text_generation_config.temperature);
|
.unwrap_or(text_generation_config.temperature);
|
||||||
|
|
||||||
let mut request_builder = CreateChatCompletionRequestArgs::default();
|
let mut request_builder = CreateResponseArgs::default();
|
||||||
|
|
||||||
request_builder
|
request_builder
|
||||||
.model(&text_generation_config.model_id)
|
.model(&text_generation_config.model_id)
|
||||||
.temperature(temperature)
|
.temperature(temperature)
|
||||||
.messages(openai_conversation_messages);
|
.input(input);
|
||||||
|
|
||||||
if let Some(max_response_tokens) = text_generation_config.max_response_tokens {
|
let mut tools = Vec::new();
|
||||||
request_builder.max_tokens(max_response_tokens);
|
if text_generation_config.tools.web_search {
|
||||||
|
tools.push(Tool::WebSearch(WebSearchTool::default()));
|
||||||
|
}
|
||||||
|
if text_generation_config.tools.code_interpreter {
|
||||||
|
tools.push(Tool::CodeInterpreter(CodeInterpreterTool {
|
||||||
|
container: CodeInterpreterToolContainer::Auto(
|
||||||
|
CodeInterpreterContainerAuto::default(),
|
||||||
|
),
|
||||||
|
}));
|
||||||
}
|
}
|
||||||
|
|
||||||
if let Some(max_completion_tokens) = text_generation_config.max_completion_tokens {
|
if !tools.is_empty() {
|
||||||
request_builder.max_completion_tokens(max_completion_tokens);
|
request_builder.tools(tools);
|
||||||
|
}
|
||||||
|
|
||||||
|
if let Some(max_response_tokens) = text_generation_config.max_response_tokens {
|
||||||
|
request_builder.max_output_tokens(max_response_tokens);
|
||||||
|
} else if let Some(max_completion_tokens) = text_generation_config.max_completion_tokens {
|
||||||
|
request_builder.max_output_tokens(max_completion_tokens);
|
||||||
}
|
}
|
||||||
|
|
||||||
let request = request_builder.build()?;
|
let request = request_builder.build()?;
|
||||||
@@ -160,33 +179,31 @@ impl ControllerTrait for Controller {
|
|||||||
model = format!("{:?}", request.model),
|
model = format!("{:?}", request.model),
|
||||||
?messages_count,
|
?messages_count,
|
||||||
request = request_as_json,
|
request = request_as_json,
|
||||||
"Sending OpenAI chat completion API request"
|
"Sending OpenAI response API request"
|
||||||
);
|
);
|
||||||
}
|
}
|
||||||
|
|
||||||
let response = self.client.chat().create(request).await?;
|
let response = self.client.responses().create(request).await?;
|
||||||
|
|
||||||
tracing::trace!(
|
tracing::trace!(
|
||||||
?response,
|
?response,
|
||||||
"Got response from the OpenAI chat completion API"
|
"Got response from the OpenAI response API"
|
||||||
);
|
);
|
||||||
|
|
||||||
// We only request 1 result, so there should only be 1 choice.
|
for item in response.output {
|
||||||
if let Some(choice) = response.choices.into_iter().next() {
|
if let OutputItem::Message(message) = item {
|
||||||
match choice.message.content {
|
for content in message.content {
|
||||||
Some(text) => {
|
if let OutputMessageContent::OutputText(text_content) = content {
|
||||||
return Ok(TextGenerationResult { text });
|
return Ok(TextGenerationResult {
|
||||||
}
|
text: text_content.text,
|
||||||
None => {
|
});
|
||||||
return Err(anyhow::anyhow!(
|
}
|
||||||
"No content was found in the response choice from the OpenAI chat completion API"
|
|
||||||
));
|
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
Err(anyhow::anyhow!(
|
Err(anyhow::anyhow!(
|
||||||
"No response messages choices were returned from the OpenAI chat completion API"
|
"No response messages choices were returned from the OpenAI response API"
|
||||||
))
|
))
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -1,11 +1,7 @@
|
|||||||
use async_openai::types::{
|
use async_openai::types::{
|
||||||
chat::{
|
responses::{
|
||||||
ChatCompletionRequestAssistantMessageArgs, ChatCompletionRequestMessage,
|
EasyInputContent, EasyInputMessage, ImageDetail, InputContent, InputImageContent,
|
||||||
ChatCompletionRequestMessageContentPartImage,
|
InputItem, InputParam, MessageType, Role,
|
||||||
ChatCompletionRequestSystemMessageArgs,
|
|
||||||
ChatCompletionRequestUserMessageArgs, ChatCompletionRequestUserMessageContent,
|
|
||||||
ChatCompletionRequestUserMessageContentPart,
|
|
||||||
ImageUrlArgs,
|
|
||||||
},
|
},
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -14,79 +10,43 @@ use crate::conversation::llm::{
|
|||||||
};
|
};
|
||||||
use crate::utils::base64::base64_encode;
|
use crate::utils::base64::base64_encode;
|
||||||
|
|
||||||
pub fn convert_llm_messages_to_openai_messages(
|
pub fn convert_llm_messages_to_openai_response_input(
|
||||||
conversation_messages: Vec<LLMMessage>,
|
conversation_messages: Vec<LLMMessage>,
|
||||||
) -> Vec<ChatCompletionRequestMessage> {
|
) -> InputParam {
|
||||||
let mut openai_conversation_messages: Vec<ChatCompletionRequestMessage> =
|
let mut items = Vec::with_capacity(conversation_messages.len());
|
||||||
Vec::with_capacity(conversation_messages.len());
|
|
||||||
|
|
||||||
for message in conversation_messages {
|
for message in conversation_messages {
|
||||||
let openai_message = convert_llm_message_to_openai_message(message);
|
let role = match message.author {
|
||||||
if let Some(openai_message) = openai_message {
|
LLMAuthor::Prompt => Role::System,
|
||||||
openai_conversation_messages.push(openai_message);
|
LLMAuthor::Assistant => Role::Assistant,
|
||||||
}
|
LLMAuthor::User => Role::User,
|
||||||
}
|
};
|
||||||
|
|
||||||
openai_conversation_messages
|
let content = match message.content {
|
||||||
}
|
LLMMessageContent::Text(text) => EasyInputContent::Text(text),
|
||||||
|
LLMMessageContent::Image(image_details) => {
|
||||||
|
let image_url = format!(
|
||||||
|
"data:{};base64,{}",
|
||||||
|
image_details.mime,
|
||||||
|
base64_encode(&image_details.data)
|
||||||
|
);
|
||||||
|
|
||||||
fn convert_llm_message_to_openai_message(
|
EasyInputContent::ContentList(vec![InputContent::InputImage(InputImageContent {
|
||||||
llm_message: LLMMessage,
|
image_url: Some(image_url),
|
||||||
) -> Option<ChatCompletionRequestMessage> {
|
detail: ImageDetail::Auto,
|
||||||
match &llm_message.content {
|
file_id: None,
|
||||||
LLMMessageContent::Text(text) => Some(match llm_message.author {
|
})])
|
||||||
LLMAuthor::Prompt => ChatCompletionRequestSystemMessageArgs::default()
|
|
||||||
.content(text.clone())
|
|
||||||
.build()
|
|
||||||
.expect("Failed building OpenAI system message")
|
|
||||||
.into(),
|
|
||||||
LLMAuthor::Assistant => ChatCompletionRequestAssistantMessageArgs::default()
|
|
||||||
.content(text.clone())
|
|
||||||
.build()
|
|
||||||
.expect("Failed building OpenAI assistant message")
|
|
||||||
.into(),
|
|
||||||
LLMAuthor::User => ChatCompletionRequestUserMessageArgs::default()
|
|
||||||
.content(text.clone())
|
|
||||||
.build()
|
|
||||||
.expect("Failed building OpenAI user message")
|
|
||||||
.into(),
|
|
||||||
}),
|
|
||||||
LLMMessageContent::Image(image_details) => {
|
|
||||||
let image_url = format!(
|
|
||||||
"data:{};base64,{}",
|
|
||||||
image_details.mime,
|
|
||||||
base64_encode(&image_details.data)
|
|
||||||
);
|
|
||||||
|
|
||||||
let part = ChatCompletionRequestUserMessageContentPart::ImageUrl(
|
|
||||||
ChatCompletionRequestMessageContentPartImage {
|
|
||||||
image_url: ImageUrlArgs::default()
|
|
||||||
.url(image_url)
|
|
||||||
.build()
|
|
||||||
.expect("Failed building OpenAI image url"),
|
|
||||||
},
|
|
||||||
);
|
|
||||||
|
|
||||||
let message_content = ChatCompletionRequestUserMessageContent::Array(vec![part]);
|
|
||||||
|
|
||||||
match llm_message.author {
|
|
||||||
LLMAuthor::User => Some(
|
|
||||||
ChatCompletionRequestUserMessageArgs::default()
|
|
||||||
.content(message_content)
|
|
||||||
.build()
|
|
||||||
.expect("Failed building OpenAI user message")
|
|
||||||
.into(),
|
|
||||||
),
|
|
||||||
_ => {
|
|
||||||
tracing::warn!(
|
|
||||||
"OpenAI API does not support image content for messages authored by {:?}. This message part will be skipped.",
|
|
||||||
llm_message.author
|
|
||||||
);
|
|
||||||
None
|
|
||||||
}
|
|
||||||
}
|
}
|
||||||
}
|
};
|
||||||
|
|
||||||
|
items.push(InputItem::EasyMessage(EasyInputMessage {
|
||||||
|
r#type: MessageType::Message,
|
||||||
|
role,
|
||||||
|
content,
|
||||||
|
}));
|
||||||
}
|
}
|
||||||
|
|
||||||
|
InputParam::Items(items)
|
||||||
}
|
}
|
||||||
|
|
||||||
pub(super) fn convert_string_to_enum<T>(value: &str) -> Result<T, String>
|
pub(super) fn convert_string_to_enum<T>(value: &str) -> Result<T, String>
|
||||||
|
|||||||
@@ -95,6 +95,7 @@ impl TryInto<OpenAITextGenerationConfig> for TextGenerationConfig {
|
|||||||
max_response_tokens: self.max_response_tokens,
|
max_response_tokens: self.max_response_tokens,
|
||||||
max_completion_tokens: None,
|
max_completion_tokens: None,
|
||||||
max_context_tokens: self.max_context_tokens,
|
max_context_tokens: self.max_context_tokens,
|
||||||
|
tools: Default::default(),
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user