artificial_openai/api_v1/
chat_completion.rs1use artificial_core::error::ArtificialError;
2use artificial_core::generic::{GenericFunctionSpec, GenericMessage, GenericRole};
3use artificial_core::provider::ChatCompleteParameters;
4use serde::de::{self, Visitor};
5use serde::{Deserialize, Deserializer, Serialize};
6
7use std::fmt;
8
9use crate::impl_builder_methods;
10use crate::model_map::map_model;
11
12use super::common;
13use super::tools::ToolCall;
14
15#[derive(Debug, Serialize, Clone)]
16pub struct ChatCompletionRequest {
17 pub model: String,
18 pub messages: Vec<ChatCompletionMessage>,
19 #[serde(skip_serializing_if = "Option::is_none")]
20 pub tools: Option<Vec<ToolSpec>>,
21 #[serde(skip_serializing_if = "Option::is_none")]
22 pub temperature: Option<f64>,
23 #[serde(skip_serializing_if = "Option::is_none")]
24 pub top_p: Option<f64>,
25 #[serde(skip_serializing_if = "Option::is_none")]
26 pub n: Option<i64>,
27 #[serde(skip_serializing_if = "Option::is_none")]
28 pub response_format: Option<serde_json::Value>,
29 #[serde(skip_serializing_if = "Option::is_none")]
30 pub stream: Option<bool>,
31 #[serde(skip_serializing_if = "Option::is_none")]
32 pub tool_choice: Option<ToolChoice>,
33}
34
35impl ChatCompletionRequest {
36 pub fn new(model: String, messages: Vec<ChatCompletionMessage>) -> Self {
37 Self {
38 model,
39 messages,
40 temperature: None,
41 top_p: None,
42 n: None,
43 response_format: None,
44 stream: None,
45 tools: None,
46 tool_choice: None,
47 }
48 }
49}
50
51impl_builder_methods!(
52 ChatCompletionRequest,
53 response_format: serde_json::Value
54);
55
56impl<M> TryFrom<ChatCompleteParameters<M>> for ChatCompletionRequest
57where
58 M: Into<ChatCompletionMessage> + Clone,
59{
60 type Error = ArtificialError;
61
62 fn try_from(value: ChatCompleteParameters<M>) -> Result<Self, Self::Error> {
63 Ok(Self {
64 model: map_model(&value.model)
65 .ok_or(ArtificialError::InvalidRequest(format!(
66 "backend does not support selected model: {:?}",
67 value.model
68 )))?
69 .into(),
70 messages: value.messages.into_iter().map(Into::into).collect(),
71 tools: value
72 .tools
73 .map(|tools| tools.into_iter().map(Into::into).collect()),
74 temperature: value.temperature,
75 top_p: None,
76 n: None,
77 response_format: value.response_format,
78 stream: None,
79 tool_choice: None,
80 })
81 }
82}
83
84#[derive(Debug, Deserialize, Serialize, Clone)]
85#[serde(rename_all = "snake_case")]
86pub struct ToolSpec {
87 pub function: ToolFunctionSpec,
88 pub r#type: ToolType,
89}
90
91impl From<GenericFunctionSpec> for ToolSpec {
92 fn from(value: GenericFunctionSpec) -> Self {
93 ToolSpec {
94 function: ToolFunctionSpec {
95 name: value.name,
96 description: value.description,
97 parameters: value.parameters,
98 strict: Some(true),
99 },
100 r#type: ToolType::Function,
101 }
102 }
103}
104
105#[derive(Debug, Deserialize, Serialize, Clone)]
106#[serde(rename_all = "snake_case")]
107pub struct ToolFunctionSpec {
108 pub name: String,
109 pub description: String,
110 pub parameters: serde_json::Value,
111 pub strict: Option<bool>,
112}
113
114#[derive(Debug, Deserialize, Serialize, Copy, Clone)]
115#[serde(rename_all = "snake_case")]
116pub enum ToolType {
117 Function,
118}
119
120#[derive(Debug, Deserialize, Serialize, Copy, Clone)]
121#[serde(rename_all = "snake_case")]
122pub enum ToolChoice {
123 None,
124 Auto,
125}
126
127#[derive(Debug, Deserialize, Serialize, Clone, PartialEq, Eq)]
128#[serde(rename_all = "snake_case")]
129pub enum MessageRole {
130 User,
131 System,
132 Assistant,
133 Function,
134 Tool,
135}
136
137#[derive(Debug, Clone, PartialEq, Eq)]
138pub enum Content {
139 Text(String),
140}
141
142impl serde::Serialize for Content {
143 fn serialize<S>(&self, serializer: S) -> Result<S::Ok, S::Error>
144 where
145 S: serde::Serializer,
146 {
147 match *self {
148 Content::Text(ref text) => {
149 if text.is_empty() {
150 serializer.serialize_none()
151 } else {
152 serializer.serialize_str(text)
153 }
154 }
155 }
156 }
157}
158
159impl<'de> Deserialize<'de> for Content {
160 fn deserialize<D>(deserializer: D) -> Result<Content, D::Error>
161 where
162 D: Deserializer<'de>,
163 {
164 struct ContentVisitor;
165
166 impl<'de> Visitor<'de> for ContentVisitor {
167 type Value = Content;
168
169 fn expecting(&self, formatter: &mut fmt::Formatter) -> fmt::Result {
170 formatter.write_str("a valid content type")
171 }
172
173 fn visit_str<E>(self, value: &str) -> Result<Content, E>
174 where
175 E: de::Error,
176 {
177 Ok(Content::Text(value.to_string()))
178 }
179
180 fn visit_none<E>(self) -> Result<Self::Value, E>
181 where
182 E: de::Error,
183 {
184 Ok(Content::Text(String::new()))
185 }
186
187 fn visit_unit<E>(self) -> Result<Self::Value, E>
188 where
189 E: de::Error,
190 {
191 Ok(Content::Text(String::new()))
192 }
193 }
194
195 deserializer.deserialize_any(ContentVisitor)
196 }
197}
198
199#[allow(dead_code)]
200#[derive(Debug, Deserialize, Serialize, Clone, PartialEq, Eq)]
201#[serde(rename_all = "snake_case")]
202pub enum ContentType {
203 Text,
204}
205
206#[derive(Debug, Deserialize, Serialize, Clone)]
207pub struct ChatCompletionMessage {
208 pub role: MessageRole,
209 pub content: Option<Content>,
210 pub name: Option<String>,
211 pub tool_calls: Option<Vec<ToolCall>>,
212 pub tool_call_id: Option<String>,
213}
214
215#[derive(Debug, Deserialize, Clone)]
216pub struct ChatCompletionMessageForResponse {
217 pub role: MessageRole,
218 #[serde(skip_serializing_if = "Option::is_none")]
219 pub content: Option<String>,
220 #[serde(skip_serializing_if = "Option::is_none")]
221 pub reasoning_content: Option<String>,
222 #[serde(skip_serializing_if = "Option::is_none")]
223 pub tool_calls: Option<Vec<ToolCall>>,
224 #[serde(skip_serializing_if = "Option::is_none")]
225 pub tool_call_id: Option<String>,
226 #[serde(skip_serializing_if = "Option::is_none")]
227 pub name: Option<String>,
228}
229
230impl From<ChatCompletionMessageForResponse> for GenericMessage {
231 fn from(val: ChatCompletionMessageForResponse) -> Self {
232 GenericMessage {
233 content: val.content,
234 role: val.role.into(),
235 tool_calls: val
236 .tool_calls
237 .map(|calls| calls.into_iter().map(Into::into).collect()),
238 name: val.name,
239 tool_call_id: val.tool_call_id,
240 }
241 }
242}
243
244#[allow(dead_code)]
245#[derive(Debug, Deserialize)]
246pub struct ChatCompletionChoice {
247 pub index: i64,
248 pub message: ChatCompletionMessageForResponse,
249 pub finish_reason: Option<FinishReason>,
250 pub finish_details: Option<FinishDetails>,
251}
252
253#[allow(dead_code)]
254#[derive(Debug, Deserialize)]
255pub struct ChatCompletionResponse {
256 pub id: Option<String>,
257 pub object: String,
258 pub created: i64,
259 pub model: String,
260 pub choices: Vec<ChatCompletionChoice>,
261 pub usage: common::Usage,
262 pub system_fingerprint: Option<String>,
263}
264
265#[derive(Debug, Deserialize, PartialEq, Eq)]
266#[serde(rename_all = "snake_case")]
267pub enum FinishReason {
268 Stop,
269 Length,
270 ContentFilter,
271 ToolCalls,
272}
273
274#[allow(non_camel_case_types, dead_code)]
275#[derive(Debug, Deserialize)]
276pub struct FinishDetails {
277 pub r#type: FinishReason,
278 pub stop: String,
279}
280
281impl From<GenericRole> for MessageRole {
282 fn from(value: GenericRole) -> Self {
283 match value {
284 GenericRole::System => MessageRole::System,
285 GenericRole::Assistant => MessageRole::Assistant,
286 GenericRole::User => MessageRole::User,
287 GenericRole::Tool => MessageRole::Tool,
288 }
289 }
290}
291
292impl From<MessageRole> for GenericRole {
293 fn from(val: MessageRole) -> Self {
294 match val {
295 MessageRole::User => GenericRole::User,
296 MessageRole::System => GenericRole::System,
297 MessageRole::Assistant => GenericRole::Assistant,
298 MessageRole::Function => GenericRole::Tool,
299 MessageRole::Tool => GenericRole::Tool,
300 }
301 }
302}
303
304impl From<GenericMessage> for ChatCompletionMessage {
305 fn from(value: GenericMessage) -> Self {
306 Self {
307 role: value.role.into(),
308 content: value.content.map(Content::Text),
309 name: value.name,
310 tool_calls: value
311 .tool_calls
312 .map(|v| v.into_iter().map(Into::into).collect()),
313 tool_call_id: value.tool_call_id,
314 }
315 }
316}