add image api and audio api

This commit is contained in:
jhqxxx
2025-12-25 20:25:52 +08:00
parent bf53d3d869
commit 9a4fa636be
16 changed files with 694 additions and 72 deletions
+38
View File
@@ -104,3 +104,41 @@ pub(crate) async fn chat(
}
}
}
#[post("/remove_background", data = "<req>")]
pub(crate) async fn remove_background(req: Json<ChatCompletionParameters>) -> (Status, String) {
let response = {
let model_ref = MODEL
.get()
.cloned()
.ok_or_else(|| anyhow::anyhow!("model not init"))
.unwrap();
model_ref.write().await.generate(req.into_inner())
};
match response {
Ok(res) => {
let response_str = serde_json::to_string(&res).unwrap();
(Status::Ok, response_str)
}
Err(e) => (Status::InternalServerError, e.to_string()),
}
}
#[post("/speech", data = "<req>")]
pub(crate) async fn speech(req: Json<ChatCompletionParameters>) -> (Status, String) {
let response = {
let model_ref = MODEL
.get()
.cloned()
.ok_or_else(|| anyhow::anyhow!("model not init"))
.unwrap();
model_ref.write().await.generate(req.into_inner())
};
match response {
Ok(res) => {
let response_str = serde_json::to_string(&res).unwrap();
(Status::Ok, response_str)
}
Err(e) => (Status::InternalServerError, e.to_string()),
}
}
+8 -9
View File
@@ -1,8 +1,7 @@
use std::time::Duration;
use aha::models::WhichModel;
use aha::{models::WhichModel, utils::get_default_save_dir};
use clap::Parser;
use dirs::home_dir;
use modelscope::ModelScope;
use rocket::{
Config,
@@ -65,13 +64,6 @@ async fn download_model(model_id: &str, save_dir: &str, max_retries: u32) -> any
}
}
fn get_default_save_dir() -> Option<String> {
home_dir().map(|mut path| {
path.push(".aha"); // 在 home 目录下创建 .aha 文件夹
path.to_string_lossy().to_string()
})
}
#[tokio::main]
async fn main() -> anyhow::Result<()> {
let args = Args::parse();
@@ -86,6 +78,9 @@ async fn main() -> anyhow::Result<()> {
WhichModel::DeepSeekOCR => "deepseek-ai/DeepSeek-OCR",
WhichModel::HunyuanOCR => "Tencent-Hunyuan/HunyuanOCR",
WhichModel::PaddleOCRVL => "PaddlePaddle/PaddleOCR-VL",
WhichModel::RMBG2_0 => "AI-ModelScope/RMBG-2.0",
WhichModel::VoxCPM => "OpenBMB/VoxCPM-0.5B",
WhichModel::VoxCPM1_5 => "OpenBMB/VoxCPM1.5",
};
let model_path = match args.weight_path {
Some(path) => path,
@@ -118,6 +113,10 @@ pub async fn start_http_server(port: u16) -> anyhow::Result<()> {
});
builder = builder.mount("/chat", routes![api::chat]);
// /images/remove_background
builder = builder.mount("/images", routes![api::remove_background]);
// /images/speech
builder = builder.mount("/audio", routes![api::speech]);
builder.launch().await?;
Ok(())
+26 -1
View File
@@ -18,7 +18,8 @@ use crate::models::{
deepseek_ocr::generate::DeepseekOCRGenerateModel,
hunyuan_ocr::generate::HunyuanOCRGenerateModel, minicpm4::generate::MiniCPMGenerateModel,
paddleocr_vl::generate::PaddleOCRVLGenerateModel, qwen2_5vl::generate::Qwen2_5VLGenerateModel,
qwen3vl::generate::Qwen3VLGenerateModel,
qwen3vl::generate::Qwen3VLGenerateModel, rmbg2_0::generate::RMBG2_0Model,
voxcpm::generate::VoxCPMGenerate,
};
#[derive(Debug, Clone, Copy, PartialEq, Eq, clap::ValueEnum)]
@@ -43,6 +44,12 @@ pub enum WhichModel {
HunyuanOCR,
#[value(name = "paddleocr-vl")]
PaddleOCRVL,
#[value(name = "RMBG2.0")]
RMBG2_0,
#[value(name = "voxcpm")]
VoxCPM,
#[value(name = "voxcpm1.5")]
VoxCPM1_5,
}
pub trait GenerateModel {
@@ -67,6 +74,8 @@ pub enum ModelInstance<'a> {
DeepSeekOCR(DeepseekOCRGenerateModel),
HunyuanOCR(HunyuanOCRGenerateModel<'a>),
PaddleOCRVL(Box<PaddleOCRVLGenerateModel<'a>>),
RMBG2_0(Box<RMBG2_0Model>),
VoxCPM(Box<VoxCPMGenerate>),
}
impl<'a> GenerateModel for ModelInstance<'a> {
@@ -78,6 +87,8 @@ impl<'a> GenerateModel for ModelInstance<'a> {
ModelInstance::DeepSeekOCR(model) => model.generate(mes),
ModelInstance::HunyuanOCR(model) => model.generate(mes),
ModelInstance::PaddleOCRVL(model) => model.generate(mes),
ModelInstance::RMBG2_0(model) => model.generate(mes),
ModelInstance::VoxCPM(model) => model.generate(mes),
}
}
@@ -99,6 +110,8 @@ impl<'a> GenerateModel for ModelInstance<'a> {
ModelInstance::DeepSeekOCR(model) => model.generate_stream(mes),
ModelInstance::HunyuanOCR(model) => model.generate_stream(mes),
ModelInstance::PaddleOCRVL(model) => model.generate_stream(mes),
ModelInstance::RMBG2_0(model) => model.generate_stream(mes),
ModelInstance::VoxCPM(model) => model.generate_stream(mes),
}
}
}
@@ -145,6 +158,18 @@ pub fn load_model(model_type: WhichModel, path: &str) -> Result<ModelInstance<'_
let model = PaddleOCRVLGenerateModel::init(path, None, None)?;
ModelInstance::PaddleOCRVL(Box::new(model))
}
WhichModel::RMBG2_0 => {
let model = RMBG2_0Model::init(path, None, None)?;
ModelInstance::RMBG2_0(Box::new(model))
}
WhichModel::VoxCPM => {
let model = VoxCPMGenerate::init(path, None, None)?;
ModelInstance::VoxCPM(Box::new(model))
}
WhichModel::VoxCPM1_5 => {
let model = VoxCPMGenerate::init(path, None, None)?;
ModelInstance::VoxCPM(Box::new(model))
}
};
Ok(model)
}
+50 -6
View File
@@ -1,18 +1,24 @@
use aha_openai_dive::v1::resources::chat::ChatCompletionParameters;
use std::io::Cursor;
use aha_openai_dive::v1::resources::chat::{
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
};
use anyhow::Result;
use base64::{Engine, prelude::BASE64_STANDARD};
use candle_core::{DType, Device, Tensor};
use candle_nn::VarBuilder;
use image::{Rgba, RgbaImage};
use rocket::futures::{Stream, stream};
use crate::{
models::rmbg2_0::model::BiRefNet,
models::{GenerateModel, rmbg2_0::model::BiRefNet},
utils::{
find_type_files, get_device, get_dtype,
build_img_completion_response, find_type_files, get_device, get_dtype,
img_utils::{extract_images, float_tensor_to_dynamic_image, img_transform_with_resize},
},
};
pub struct RMBG2_0 {
pub struct RMBG2_0Model {
model: BiRefNet,
h: u32,
w: u32,
@@ -20,9 +26,10 @@ pub struct RMBG2_0 {
img_std: Tensor,
device: Device,
dtype: DType,
model_name: String,
}
impl RMBG2_0 {
impl RMBG2_0Model {
pub fn init(path: &str, device: Option<&Device>, dtype: Option<DType>) -> Result<Self> {
let device = get_device(device);
let dtype = get_dtype(dtype, "float32");
@@ -41,10 +48,11 @@ impl RMBG2_0 {
img_std,
device,
dtype,
model_name: "RMBG2.0".to_string(),
})
}
pub fn generate(&self, mes: ChatCompletionParameters) -> Result<Vec<RgbaImage>> {
pub fn inference(&self, mes: ChatCompletionParameters) -> Result<Vec<RgbaImage>> {
let imgs = extract_images(&mes)?;
let mut rmbg_png = vec![];
for img in imgs {
@@ -78,3 +86,39 @@ impl RMBG2_0 {
Ok(rmbg_png)
}
}
impl GenerateModel for RMBG2_0Model {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let rmbg_png = self.inference(mes)?;
let mut base64_vec = vec![];
for img in rmbg_png {
let mut png_bytes = Vec::new();
img.write_to(&mut Cursor::new(&mut png_bytes), image::ImageFormat::Png)?;
let base64_string = BASE64_STANDARD.encode(png_bytes);
base64_vec.push(base64_string);
}
let response = build_img_completion_response(&base64_vec, &self.model_name);
Ok(response)
}
#[allow(unused_variables)]
fn generate_stream(
&mut self,
mes: ChatCompletionParameters,
) -> Result<
Box<
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
+ Send
+ Unpin
+ '_,
>,
> {
let error_stream = stream::once(async {
Err(anyhow::anyhow!(format!(
"{} model not support stream",
self.model_name
))) as Result<ChatCompletionChunkResponse, anyhow::Error>
});
Ok(Box::new(Box::pin(error_stream)))
}
}
+86 -9
View File
@@ -1,22 +1,36 @@
use std::collections::HashMap;
use aha_openai_dive::v1::resources::chat::{
ChatCompletionChunkResponse, ChatCompletionParameters, ChatCompletionResponse,
};
use anyhow::{Ok, Result};
use base64::{Engine, prelude::BASE64_STANDARD};
use candle_core::{DType, Device, Tensor, pickle::read_all_with_key};
use candle_nn::VarBuilder;
use rocket::futures::{Stream, stream};
use crate::{
models::voxcpm::{
audio_vae::AudioVAE,
config::{AudioVaeConfig, VoxCPMConfig},
model::VoxCPMModel,
tokenizer::SingleChineseTokenizer,
models::{
GenerateModel,
voxcpm::{
audio_vae::AudioVAE,
config::{AudioVaeConfig, VoxCPMConfig},
model::VoxCPMModel,
tokenizer::SingleChineseTokenizer,
},
},
utils::{
audio_utils::{extract_audio_url, get_audio_wav_u8},
build_audio_completion_response, extract_metadata_value, extract_user_text,
find_type_files, get_device, get_dtype,
},
utils::{find_type_files, get_device, get_dtype},
};
pub struct VoxCPMGenerate {
voxcpm: VoxCPMModel,
prompt_cache: Option<HashMap<String, Tensor>>,
sample_rate: usize,
model_name: String,
}
impl VoxCPMGenerate {
@@ -48,6 +62,11 @@ impl VoxCPMGenerate {
sample_rate: 16000,
},
};
let model_name = if audio_config.sample_rate == 16000 {
"VoxCPM".to_string()
} else {
"VoxCPM1.5".to_string()
};
let audio_vae = AudioVAE::new(
vb_vae,
audio_config.encoder_dim,
@@ -85,6 +104,8 @@ impl VoxCPMGenerate {
Ok(Self {
voxcpm,
prompt_cache: None,
sample_rate: audio_config.sample_rate,
model_name,
})
}
@@ -135,7 +156,7 @@ impl VoxCPMGenerate {
prompt_text: Option<String>,
prompt_wav_path: Option<String>,
) -> Result<Tensor> {
let audio = self.generate(
let audio = self.inference(
target_text,
prompt_text,
prompt_wav_path,
@@ -150,10 +171,10 @@ impl VoxCPMGenerate {
}
pub fn generate_simple(&mut self, target_text: String) -> Result<Tensor> {
// let audio = self.generate(target_text, None, None, 2, 100, 10, 2.0, false, 6.0)?;
let audio = self.generate(target_text, None, None, 2, 100, 10, 2.0, 6.0)?;
let audio = self.inference(target_text, None, None, 2, 100, 10, 2.0, 6.0)?;
Ok(audio)
}
pub fn generate(
pub fn inference(
&mut self,
target_text: String,
prompt_text: Option<String>,
@@ -179,3 +200,59 @@ impl VoxCPMGenerate {
Ok(audio)
}
}
impl GenerateModel for VoxCPMGenerate {
fn generate(&mut self, mes: ChatCompletionParameters) -> Result<ChatCompletionResponse> {
let prompt_text = extract_metadata_value::<String>(&mes.metadata, "prompt_text");
let min_len = extract_metadata_value::<usize>(&mes.metadata, "min_len").unwrap_or(2);
let max_len = extract_metadata_value::<usize>(&mes.metadata, "max_len").unwrap_or(4096);
let inference_timesteps =
extract_metadata_value::<usize>(&mes.metadata, "inference_timesteps").unwrap_or(10);
let cfg_value = extract_metadata_value::<f64>(&mes.metadata, "cfg_value").unwrap_or(2.0);
let retry_badcase_ratio_threshold =
extract_metadata_value::<f64>(&mes.metadata, "retry_badcase_ratio_threshold")
.unwrap_or(6.0);
let target_text = extract_user_text(&mes)?;
let prompt_wav = extract_audio_url(&mes)?;
let prompt_wav_path = if !prompt_wav.is_empty() {
Some(prompt_wav[0].clone())
} else {
None
};
let audio = self.voxcpm.generate(
target_text,
prompt_text,
prompt_wav_path,
min_len,
max_len,
inference_timesteps,
cfg_value,
retry_badcase_ratio_threshold,
)?;
let wav_u8 = get_audio_wav_u8(&audio, self.sample_rate as u32)?;
let base64_audio = BASE64_STANDARD.encode(wav_u8);
let response = build_audio_completion_response(&base64_audio, &self.model_name);
Ok(response)
}
#[allow(unused_variables)]
fn generate_stream(
&mut self,
mes: ChatCompletionParameters,
) -> Result<
Box<
dyn Stream<Item = Result<ChatCompletionChunkResponse, anyhow::Error>>
+ Send
+ Unpin
+ '_,
>,
> {
let error_stream = stream::once(async {
Err(anyhow::anyhow!(format!(
"{} model not support stream",
self.model_name
))) as Result<ChatCompletionChunkResponse, anyhow::Error>
});
Ok(Box::new(Box::pin(error_stream)))
}
}
+6 -3
View File
@@ -513,7 +513,7 @@ impl VoxCPMModel {
let text_token = Tensor::cat(&[text_token, audio_start], D::Minus1)?;
let text_length = text_token.dim(0)?;
let mut audio =
load_audio_with_resample(path, self.device.clone(), Some(self.sample_rate))?;
load_audio_with_resample(&path, self.device.clone(), Some(self.sample_rate))?;
let patch_len = self.patch_size * self.chunk_size;
if audio.dim(1)? % patch_len != 0 {
audio = audio.pad_with_zeros(
@@ -728,8 +728,11 @@ impl VoxCPMModel {
) -> Result<HashMap<String, Tensor>> {
let text_token = self.tokenizer.encode(prompt_text)?;
let text_token = Tensor::from_slice(&text_token, text_token.len(), &self.device)?;
let mut audio =
load_audio_with_resample(prompt_wav_path, self.device.clone(), Some(self.sample_rate))?;
let mut audio = load_audio_with_resample(
&prompt_wav_path,
self.device.clone(),
Some(self.sample_rate),
)?;
let patch_len = self.patch_size * self.chunk_size;
if audio.dim(1)? % patch_len != 0 {
audio = audio.pad_with_zeros(D::Minus1, 0, patch_len - audio.dim(1)? % patch_len)?;
+180 -6
View File
@@ -1,12 +1,22 @@
use std::f64::consts::PI;
use std::path::Path;
use std::fs::File;
use std::io::Write;
use std::path::{Path, PathBuf};
use std::{f64::consts::PI, io::Cursor};
use aha_openai_dive::v1::resources::chat::{
ChatCompletionParameters, ChatCompletionResponse, ChatMessage, ChatMessageContent,
ChatMessageContentPart,
};
use anyhow::{Result, anyhow};
use base64::Engine;
use base64::prelude::BASE64_STANDARD;
use candle_core::{D, Device, Tensor};
use candle_nn::{Conv1d, Conv1dConfig, Module};
use hound::{SampleFormat, WavReader};
use num::integer::gcd;
use crate::utils::get_default_save_dir;
// 重采样方法枚举
#[derive(Debug, Clone, Copy)]
pub enum ResamplingMethod {
@@ -213,9 +223,62 @@ pub fn resample_simple(waveform: &Tensor, orig_freq: i64, new_freq: i64) -> Resu
None,
)
}
pub fn load_audio_from_url(url: &str) -> Result<PathBuf> {
tokio::task::block_in_place(|| {
let client = reqwest::blocking::Client::new();
let response = client.get(url).send()?;
if !response.status().is_success() {
return Err(anyhow::anyhow!(
"Failed to download file: {}",
response.status()
));
}
let temp_dir = get_default_save_dir().expect("Failed to get home directory");
let temp_dir = PathBuf::from(temp_dir);
let temp_path = temp_dir.join("temp_audio.wav");
pub fn load_audio<P: AsRef<Path>>(path: P, device: Device) -> Result<(Tensor, usize)> {
let mut reader = WavReader::open(path)?;
let mut file = std::fs::File::create(&temp_path)?;
let mut content = Cursor::new(response.bytes()?);
std::io::copy(&mut content, &mut file)?;
// Return the temp directory to keep it alive until the function ends
Ok(temp_path)
})
}
pub fn get_audio_path(path_str: &str) -> Result<PathBuf> {
if path_str.starts_with("http://") || path_str.starts_with("https://") {
// Download file from network
load_audio_from_url(path_str)
} else if path_str.starts_with("file://") {
// Convert file:// URL to local path
let path = url::Url::parse(path_str)?;
let path = path.to_file_path();
let path = match path {
Ok(path) => path,
Err(_) => {
let mut path = path_str.to_owned();
path = path.split_off(7);
PathBuf::from(path)
}
};
Ok(path)
} else if path_str.starts_with("data:audio") && path_str.contains("base64,") {
let data: Vec<&str> = path_str.split("base64,").collect();
let data = data[1];
let temp_dir = get_default_save_dir().expect("Failed to get home directory");
let temp_dir = PathBuf::from(temp_dir);
let temp_path = temp_dir.join("temp_audio.wav");
save_audio_from_base64(data, &temp_path)?;
Ok(temp_path)
} else {
Err(anyhow::anyhow!("get audio path error {}", path_str))
}
}
pub fn load_audio(path: &str, device: Device) -> Result<(Tensor, usize)> {
let audio_path = get_audio_path(path)?;
let mut reader = WavReader::open(audio_path)?;
let spec = reader.spec();
let samples: Vec<f32> = match spec.sample_format {
SampleFormat::Int => {
@@ -264,8 +327,8 @@ pub fn load_audio<P: AsRef<Path>>(path: P, device: Device) -> Result<(Tensor, us
Ok((audio_tensor, sample_rate as usize))
}
pub fn load_audio_with_resample<P: AsRef<Path>>(
path: P,
pub fn load_audio_with_resample(
path: &str,
device: Device,
target_sample_rate: Option<usize>,
) -> Result<Tensor> {
@@ -299,3 +362,114 @@ pub fn save_wav(audio: &Tensor, save_path: &str, sample_rate: u32) -> Result<()>
writer.finalize().unwrap();
Ok(())
}
pub fn get_audio_wav_u8(audio: &Tensor, sample_rate: u32) -> Result<Vec<u8>> {
let spec = hound::WavSpec {
channels: 1,
sample_rate,
bits_per_sample: 16,
sample_format: hound::SampleFormat::Int,
};
assert_eq!(audio.dim(0)?, 1, "audio channel must be 1");
let max = audio.abs()?.max_all()?;
let max = max.to_scalar::<f32>()?;
let ratio = if max > 1.0 { 32767.0 / max } else { 32767.0 };
let audio = audio.squeeze(0)?;
let audio_vec = audio.to_vec1::<f32>()?;
let mut cursor = Cursor::new(Vec::new());
let mut writer = hound::WavWriter::new(&mut cursor, spec)?;
for i in audio_vec {
let sample_i16 = (i * ratio).round() as i16;
writer.write_sample(sample_i16)?;
}
writer.finalize()?;
let wav_buffer = cursor.into_inner();
Ok(wav_buffer)
}
pub fn extract_audio_url(mes: &ChatCompletionParameters) -> Result<Vec<String>> {
let mut audio_vec = Vec::new();
for chat_mes in mes.messages.clone() {
if let ChatMessage::User { content, .. } = chat_mes.clone()
&& let ChatMessageContent::ContentPart(part_vec) = content
{
for part in part_vec {
if let ChatMessageContentPart::Audio(audio_part) = part {
let audio_url = audio_part.audio_url;
audio_vec.push(audio_url.url);
}
}
}
// if let ChatMessage::User { content, .. } = chat_mes.clone()
// && let ChatMessageContent::ContentPart(part_vec) = content
// {
// for part in part_vec {
// if let ChatMessageContentPart::Text(text_part) = part {
// let text = text_part.text;
// if text.chars().count() > 0 {
// ret = ret + &text + "\n"
// }
// }
// }
// }
}
Ok(audio_vec)
}
// 从 ChatCompletionResponse 中提取音频数据
pub fn extract_audio_base64_from_response(
response: &ChatCompletionResponse,
) -> Result<Vec<String>> {
let mut audio_data_list = Vec::new();
for choice in &response.choices {
if let ChatMessage::Assistant {
content: Some(ChatMessageContent::ContentPart(parts)),
..
} = &choice.message
{
for part in parts.clone() {
if let ChatMessageContentPart::Audio(audio_part) = part {
// if let Some(audio_data) = &audio_part.audio_url {
// audio_data_list.push(audio_data.data.clone());
// }
let audio_url = audio_part.audio_url;
audio_data_list.push(audio_url.url);
}
}
}
}
Ok(audio_data_list)
}
// 将 base64 音频数据解码并保存到文件
pub fn save_audio_from_base64<P: AsRef<Path>>(base64_data: &str, file_path: P) -> Result<()> {
// 解码 base64 数据
let data: Vec<&str> = base64_data.split("base64,").collect();
let data = data[1];
let decoded_data = BASE64_STANDARD.decode(data)?;
// 创建文件并写入数据
let mut file = File::create(file_path)?;
file.write_all(&decoded_data)?;
Ok(())
}
// 组合函数:从响应中提取音频并保存到文件
pub fn extract_and_save_audio_from_response(
response: &ChatCompletionResponse,
directory: &str,
) -> Result<Vec<String>> {
let audio_data_list = extract_audio_base64_from_response(response)?;
let mut saved_files = Vec::new();
for (index, audio_data) in audio_data_list.iter().enumerate() {
let file_path = format!("{}/audio_{}.wav", directory, index);
save_audio_from_base64(audio_data, &file_path)?;
saved_files.push(file_path);
}
Ok(saved_files)
}
+136 -8
View File
@@ -3,19 +3,21 @@ pub mod img_utils;
pub mod tensor_utils;
pub mod video_utils;
use std::process::Command;
use std::{fs, process::Command};
use aha_openai_dive::v1::resources::{
chat::{
ChatCompletionChoice, ChatCompletionChunkChoice, ChatCompletionChunkResponse,
ChatCompletionParameters, ChatCompletionResponse, ChatMessage, ChatMessageContent,
ChatMessageContentPart, DeltaChatMessage, DeltaFunction, DeltaToolCall, Function, ToolCall,
AudioUrlType, ChatCompletionChoice, ChatCompletionChunkChoice, ChatCompletionChunkResponse,
ChatCompletionParameters, ChatCompletionResponse, ChatMessage, ChatMessageAudioContentPart,
ChatMessageContent, ChatMessageContentPart, ChatMessageImageContentPart, DeltaChatMessage,
DeltaFunction, DeltaToolCall, Function, ImageUrlType, ToolCall,
},
shared::{FinishReason, Usage},
};
use anyhow::Result;
use candle_core::{DType, Device};
use candle_transformers::generation::{LogitsProcessor, Sampling};
use dirs::home_dir;
pub fn get_device(device: Option<&Device>) -> Device {
match device {
@@ -137,6 +139,90 @@ pub fn ceil_by_factor(num: f32, factor: u32) -> u32 {
ceil * factor
}
pub fn build_img_completion_response(
base64vec: &Vec<String>,
model_name: &str,
) -> ChatCompletionResponse {
let id = uuid::Uuid::new_v4().to_string();
let mut response = ChatCompletionResponse {
id: Some(id),
choices: vec![],
created: chrono::Utc::now().timestamp() as u32,
model: model_name.to_string(),
service_tier: None,
system_fingerprint: None,
object: "chat.completion".to_string(),
usage: None,
};
let mut conten_part_vec = vec![];
for img_bas64 in base64vec {
let img_base64_prefix = "data:image/png;base64,".to_string() + img_bas64;
let part = ChatMessageContentPart::Image(ChatMessageImageContentPart {
r#type: "image".to_string(),
image_url: ImageUrlType {
url: img_base64_prefix,
detail: None,
},
});
conten_part_vec.push(part);
}
let choice = ChatCompletionChoice {
index: 0,
message: ChatMessage::Assistant {
content: Some(ChatMessageContent::ContentPart(conten_part_vec)),
reasoning_content: None,
refusal: None,
name: None,
audio: None,
tool_calls: None,
},
finish_reason: Some(FinishReason::StopSequenceReached),
logprobs: None,
};
response.choices.push(choice);
response
}
pub fn build_audio_completion_response(
base64_audio: &String,
model_name: &str,
) -> ChatCompletionResponse {
let id = uuid::Uuid::new_v4().to_string();
let mut response = ChatCompletionResponse {
id: Some(id),
choices: vec![],
created: chrono::Utc::now().timestamp() as u32,
model: model_name.to_string(),
service_tier: None,
system_fingerprint: None,
object: "chat.completion".to_string(),
usage: None,
};
let base64_audio = format!("data:audio/wav;base64,{}", base64_audio);
let conten_part_vec = vec![ChatMessageContentPart::Audio(ChatMessageAudioContentPart {
r#type: "audio".to_string(),
audio_url: AudioUrlType {
url: base64_audio.to_string(),
},
})];
let choice = ChatCompletionChoice {
index: 0,
message: ChatMessage::Assistant {
content: Some(ChatMessageContent::ContentPart(conten_part_vec)),
reasoning_content: None,
refusal: None,
name: None,
audio: None,
tool_calls: None,
},
finish_reason: Some(FinishReason::StopSequenceReached),
logprobs: None,
};
response.choices.push(choice);
response
}
pub fn build_completion_response(
res: String,
model_name: &str,
@@ -144,10 +230,6 @@ pub fn build_completion_response(
) -> ChatCompletionResponse {
let id = uuid::Uuid::new_v4().to_string();
let usage = num_tokens.map(|num| Usage {
input_tokens: None,
input_tokens_details: None,
output_tokens: None,
output_tokens_details: None,
prompt_tokens: None,
completion_tokens: None,
total_tokens: num,
@@ -358,3 +440,49 @@ pub fn extract_mes(mes: &ChatCompletionParameters) -> Result<Vec<(String, String
}
Ok(mes_vec)
}
pub fn extract_metadata_value<T>(
metadata: &Option<std::collections::HashMap<String, String>>,
key: &str,
) -> Option<T>
where
T: std::str::FromStr + Clone + PartialEq,
{
if let Some(map) = metadata
&& let Some(value_str) = map.get(key)
&& let Ok(value) = value_str.parse::<T>()
{
return Some(value);
}
None
}
pub fn extract_user_text(mes: &ChatCompletionParameters) -> Result<String> {
let mut ret = "".to_string();
for chat_mes in mes.messages.clone() {
if let ChatMessage::User { content, .. } = chat_mes.clone()
&& let ChatMessageContent::ContentPart(part_vec) = content
{
for part in part_vec {
if let ChatMessageContentPart::Text(text_part) = part {
let text = text_part.text;
if text.chars().count() > 0 {
ret = ret + &text + "\n"
}
}
}
}
}
ret = ret.trim().to_string();
Ok(ret)
}
pub fn get_default_save_dir() -> Option<String> {
home_dir().map(|mut path| {
path.push(".aha");
if let Err(e) = fs::create_dir_all(&path) {
eprintln!("Failed to create directory {:?}: {}", path, e);
}
path.to_string_lossy().to_string()
})
}