use crate::io::{
api::{ApiResult, Configuration, Endpoint, Fallback, IntoBody, Param, Params, RemoteResource, TextResponse},
http::{HttpResponse, ReqwestHttpService},
};
use crate::param;
use crate::util::constants::app::DEFAULT_OPENAI_DOMAIN;
use crate::util::constants::env::{OPENAI_API_KEY, OPENAI_SERVER_HOST};
use acorn_core::options::{ApiExtension, ApiOptions};
use color_eyre::eyre::eyre;
use secrecy::ExposeSecret;
use serde::{Deserialize, Serialize};
use serde_json::Value;
mod inference;
pub use inference::{Client, InferenceError};
const RESPONSE_SCHEMA_NAME: &str = "acorn_response";
pub type AudioSpeechResponse = TextResponse;
pub type Options = ApiOptions<Extension, Param>;
pub type Response = Value;
enum BodyMode {
Empty,
Required,
}
#[derive(Clone, Copy, Debug, Eq, PartialEq)]
pub enum EndpointMode {
ChatCompletions,
Responses,
}
#[derive(Clone, Debug, Deserialize)]
#[serde(tag = "type", rename_all = "snake_case")]
enum ResponsesContent {
OutputText {
text: String,
},
Refusal {
refusal: String,
},
#[serde(other)]
Other,
}
#[derive(Clone, Debug, Deserialize)]
struct ChatChoice {
finish_reason: Option<String>,
message: ChatMessageResponse,
}
#[derive(Clone, Debug, Serialize)]
struct ChatCompletionRequest {
messages: Vec<ChatMessageRequest>,
model: String,
#[serde(skip_serializing_if = "Option::is_none")]
response_format: Option<ChatResponseFormat>,
store: bool,
stream: bool,
}
#[derive(Clone, Debug, Deserialize)]
struct ChatCompletionResponse {
choices: Vec<ChatChoice>,
id: String,
model: Option<String>,
usage: Option<ChatUsage>,
}
#[derive(Clone, Debug, Serialize)]
struct ChatMessageRequest {
content: String,
role: &'static str,
}
#[derive(Clone, Debug, Deserialize)]
struct ChatMessageResponse {
content: Option<String>,
refusal: Option<String>,
}
#[derive(Clone, Debug, Serialize)]
struct ChatResponseFormat {
json_schema: JsonSchemaDefinition,
#[serde(rename = "type")]
kind: &'static str,
}
#[derive(Clone, Copy, Debug, Deserialize)]
struct ChatUsage {
completion_tokens: Option<u64>,
prompt_tokens: Option<u64>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct Error {
pub code: Option<String>,
pub message: String,
pub param: Option<String>,
#[serde(rename = "type")]
pub error_type: String,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct ErrorResponse {
pub error: Error,
}
#[derive(Clone, Debug, Default)]
pub struct Extension;
#[derive(Clone, Debug, Serialize)]
struct JsonSchemaDefinition {
#[serde(rename = "type")]
kind: &'static str,
name: &'static str,
schema: Value,
strict: bool,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct ListModelsResponse {
pub object: String,
pub data: Vec<Model>,
}
#[derive(Clone, Debug, Serialize, Deserialize)]
pub struct Model {
pub id: String,
pub object: String,
pub created: i64,
pub owned_by: String,
}
#[derive(Clone, Debug, Deserialize)]
struct ProviderError {
code: Option<String>,
message: String,
}
#[derive(Clone, Debug, Deserialize)]
struct ProviderErrorResponse {
error: ProviderError,
#[serde(skip)]
status: u16,
}
#[derive(Clone, Debug, Deserialize)]
struct ResponsesIncompleteDetails {
reason: Option<String>,
}
#[derive(Clone, Debug, Deserialize)]
struct ResponsesOutput {
#[serde(default)]
content: Vec<ResponsesContent>,
}
#[derive(Clone, Debug, Serialize)]
struct ResponsesRequest {
input: String,
model: String,
store: bool,
stream: bool,
#[serde(skip_serializing_if = "Option::is_none")]
text: Option<ResponsesText>,
}
#[derive(Clone, Debug, Deserialize)]
struct ResponsesResponse {
error: Option<ProviderError>,
id: String,
incomplete_details: Option<ResponsesIncompleteDetails>,
model: Option<String>,
#[serde(default)]
output: Vec<ResponsesOutput>,
status: Option<String>,
usage: Option<ResponsesUsage>,
}
#[derive(Clone, Debug, Serialize)]
struct ResponsesText {
format: JsonSchemaDefinition,
}
#[derive(Clone, Copy, Debug, Deserialize)]
struct ResponsesUsage {
input_tokens: Option<u64>,
output_tokens: Option<u64>,
}
impl EndpointMode {
fn action(self) -> &'static str {
match self {
| Self::ChatCompletions => "chat-completion",
| Self::Responses => "response",
}
}
async fn invoke(self, options: &Options, loopback: bool, max_response_bytes: usize) -> ApiResult<HttpResponse> {
match prepare_request(options, self.action(), BodyMode::Required) {
| Err(why) => Err(why),
| Ok((endpoint, params)) => {
let body = serde_json::to_vec(¶ms.clone().into_body()).map_err(|why| eyre!(why));
let service = match loopback {
| true => ReqwestHttpService::loopback(),
| false => Ok(ReqwestHttpService::default()),
};
match (body, service) {
| (Err(why), _) | (_, Err(why)) => Err(why),
| (Ok(body), Ok(service)) => {
endpoint
.execute_resource(
&service,
self.action(),
Some(params),
Some(body),
Some("application/json"),
max_response_bytes,
)
.await
}
}
}
}
}
}
impl ApiExtension for Extension {
fn default_domain() -> String {
String::from(DEFAULT_OPENAI_DOMAIN)
}
fn env_token_var() -> &'static str {
OPENAI_API_KEY
}
fn env_domain_var() -> &'static str {
OPENAI_SERVER_HOST
}
}
impl From<Value> for JsonSchemaDefinition {
fn from(schema: Value) -> Self {
Self {
kind: "json_schema",
name: RESPONSE_SCHEMA_NAME,
schema,
strict: true,
}
}
}
pub async fn audio_speech(options: &Options) -> ApiResult<AudioSpeechResponse> {
invoke(options, "audio-speech", BodyMode::Required).await
}
pub async fn audio_transcription(options: &Options) -> ApiResult<Response> {
invoke(options, "audio-transcription", BodyMode::Required).await
}
pub async fn audio_voices(options: &Options) -> ApiResult<Response> {
invoke(options, "audio-voices", BodyMode::Empty).await
}
pub async fn chat_completion(options: &Options) -> ApiResult<Response> {
invoke(options, "chat-completion", BodyMode::Required).await
}
pub async fn completion(options: &Options) -> ApiResult<Response> {
invoke(options, "completion", BodyMode::Required).await
}
pub async fn embedding(options: &Options) -> ApiResult<Response> {
invoke(options, "embedding", BodyMode::Required).await
}
pub async fn image_edit(options: &Options) -> ApiResult<Response> {
invoke(options, "image-edit", BodyMode::Required).await
}
pub async fn image_generation(options: &Options) -> ApiResult<Response> {
invoke(options, "image-generation", BodyMode::Required).await
}
async fn invoke<R>(options: &Options, action: &str, body_mode: BodyMode) -> ApiResult<R>
where
R: for<'de> Deserialize<'de>,
{
match prepare_request(options, action, body_mode) {
| Ok((endpoint, params)) => {
let response = endpoint.invoke(action, Some(params)).await;
endpoint.handle_or::<R, Fallback<ErrorResponse>>(response)
}
| Err(why) => Err(why),
}
}
pub async fn models(options: &Options) -> ApiResult<ListModelsResponse> {
invoke(options, "models", BodyMode::Empty).await
}
fn prepare_request(options: &Options, action: &str, body_mode: BodyMode) -> ApiResult<(Endpoint, Vec<Param>)> {
Endpoint::from_template("openai::api")
.map(|endpoint| endpoint.with_domain(options.domain()))
.and_then(|endpoint| match (body_mode, &options.body) {
| (BodyMode::Required, Some(value)) if !value.is_empty() => Ok((
endpoint,
Params::new()
.with_auth(ExposeSecret::expose_secret(options.token()), None)
.with(param!(Body, value.as_str()))
.with_custom(options.params())
.build(),
)),
| (BodyMode::Required, _) => Err(eyre!(format!("OpenAI {action} request body is required"))),
| (BodyMode::Empty, _) => Ok((endpoint, Params::from_config(options).with_custom(options.params()).build())),
})
}
pub async fn response(options: &Options) -> ApiResult<Response> {
invoke(options, "response", BodyMode::Required).await
}
#[cfg(test)]
mod tests;