Skip to main content

sim_codec_chat/providers/openai/
encode.rs

1use serde_json::{Value, json};
2use sim_kernel::{CodecId, Error, Expr, Result};
3
4use crate::output_grammar::{OutputGrammarDialect, output_grammar_required, output_grammar_text};
5use crate::{is_model_request_expr, validate_chat_transcript};
6
7use super::OpenAiRequestOptions;
8use super::common::{
9    codec_error, codec_eval_to_codec, flatten_expr, list_field, map_field, marker_is_true,
10    optional_u64_field, string_field, symbol_field,
11};
12use crate::providers::model_params::attach_bridge_model_params;
13
14const RESERVED_MODEL_PARAM_FIELDS: &[&str] = &[
15    "model",
16    "stream",
17    "messages",
18    "tools",
19    "response_format",
20    "stream_options",
21];
22
23/// Encodes a model-request transcript into OpenAI chat-completion JSON.
24pub fn encode_openai_request(expr: &Expr, options: &OpenAiRequestOptions) -> Result<Vec<u8>> {
25    if !is_model_request_expr(expr) {
26        return Err(Error::Eval(
27            "openai codec expects a model-request transcript".to_owned(),
28        ));
29    }
30    validate_chat_transcript(expr)?;
31    let entries = request_entries(expr)?;
32    let mut payload = json!({
33        "model": options.model,
34        "stream": options.stream,
35        "messages": transcript_messages(expr)?,
36        "tools": if options.tools { Value::Array(Vec::new()) } else { Value::Null },
37    });
38    attach_output_grammar(entries, &mut payload)?;
39    if options.stream
40        && let Some(object) = payload.as_object_mut()
41    {
42        object.insert("stream_options".to_owned(), json!({"include_usage": true}));
43    }
44    if let Some(object) = payload.as_object_mut() {
45        attach_bridge_model_params(entries, object, RESERVED_MODEL_PARAM_FIELDS, "openai")?;
46    }
47    serde_json::to_vec(&payload)
48        .map_err(|err| Error::Eval(format!("openai codec failed to encode request: {err}")))
49}
50
51fn attach_output_grammar(entries: &[(Expr, Expr)], payload: &mut Value) -> Result<()> {
52    let Some(grammar) = output_grammar_text(entries, OutputGrammarDialect::JsonSchema)? else {
53        return Ok(());
54    };
55    let schema = serde_json::from_str::<Value>(&grammar)
56        .map_err(|err| Error::Eval(format!("openai output grammar is not json schema: {err}")))?;
57    let Some(object) = payload.as_object_mut() else {
58        return Err(Error::Eval(
59            "openai request payload must be a json object".to_owned(),
60        ));
61    };
62    object.insert(
63        "response_format".to_owned(),
64        json!({
65            "type": "json_schema",
66            "json_schema": {
67                "name": "sim_output",
68                "strict": output_grammar_required(entries)?,
69                "schema": schema,
70            }
71        }),
72    );
73    Ok(())
74}
75
76fn request_entries(expr: &Expr) -> Result<&[(Expr, Expr)]> {
77    let Expr::Map(entries) = expr else {
78        return Err(Error::Eval(
79            "openai codec expects request transcript as a map".to_owned(),
80        ));
81    };
82    Ok(entries)
83}
84
85/// Encodes a model-response transcript into OpenAI chat-completion JSON.
86pub fn encode_openai_response(expr: &Expr) -> Result<Vec<u8>> {
87    let value = response_json(expr)?;
88    serde_json::to_vec(&value)
89        .map_err(|err| Error::Eval(format!("openai codec failed to encode response: {err}")))
90}
91
92pub(in crate::providers) fn encode_openai_response_for_codec(
93    codec: CodecId,
94    expr: &Expr,
95) -> Result<String> {
96    if !marker_is_true(expr, "model-response") {
97        return Err(codec_error(
98            codec,
99            "openai codec expects a model-response transcript",
100        ));
101    }
102    validate_chat_transcript(expr).map_err(|err| codec_eval_to_codec(codec, err))?;
103    let value = response_json(expr).map_err(|err| codec_eval_to_codec(codec, err))?;
104    serde_json::to_string(&value).map_err(|err| codec_error(codec, err))
105}
106
107fn transcript_messages(expr: &Expr) -> Result<Vec<Value>> {
108    let Expr::Map(entries) = expr else {
109        return Err(Error::Eval(
110            "openai codec expects request transcript as a map".to_owned(),
111        ));
112    };
113    let mut messages = list_field(map_field(entries, "messages")?)?
114        .iter()
115        .map(message_to_json)
116        .collect::<Result<Vec<_>>>()?;
117    messages.push(json!({
118        "role": "user",
119        "content": [{
120            "type": "text",
121            "text": flatten_expr(map_field(entries, "task")?),
122        }],
123    }));
124    Ok(messages)
125}
126
127fn message_to_json(expr: &Expr) -> Result<Value> {
128    let Expr::Map(entries) = expr else {
129        return Err(Error::Eval("openai codec message must be a map".to_owned()));
130    };
131    Ok(json!({
132        "role": symbol_field(entries, "role")?,
133        "content": list_field(map_field(entries, "content")?)?
134            .iter()
135            .map(content_part_to_json)
136            .collect::<Result<Vec<_>>>()?,
137    }))
138}
139
140fn content_part_to_json(expr: &Expr) -> Result<Value> {
141    let Expr::Map(entries) = expr else {
142        return Err(Error::Eval(
143            "openai codec content part must be a map".to_owned(),
144        ));
145    };
146    match symbol_field(entries, "type")?.as_str() {
147        "text" => Ok(json!({
148            "type": "text",
149            "text": string_field(entries, "text")?,
150        })),
151        other => Err(Error::Eval(format!(
152            "openai codec does not support content part type {other}"
153        ))),
154    }
155}
156
157fn response_json(expr: &Expr) -> Result<Value> {
158    let Expr::Map(entries) = expr else {
159        return Err(Error::Eval(
160            "openai codec expects response transcript as a map".to_owned(),
161        ));
162    };
163    let model = string_field(entries, "model")?;
164    let finish_reason = symbol_field(entries, "stop-reason")?;
165    Ok(json!({
166        "id": "chatcmpl-sim",
167        "object": "chat.completion",
168        "created": 0,
169        "model": model,
170        "choices": [{
171            "index": 0,
172            "message": {
173                "role": "assistant",
174                "content": response_text(entries)?,
175            },
176            "finish_reason": finish_reason,
177        }],
178        "usage": response_usage(entries)?,
179    }))
180}
181
182fn response_text(entries: &[(Expr, Expr)]) -> Result<String> {
183    list_field(map_field(entries, "content")?)?
184        .iter()
185        .map(text_content)
186        .collect::<Result<Vec<_>>>()
187        .map(|parts| parts.join(""))
188}
189
190fn text_content(expr: &Expr) -> Result<String> {
191    let Expr::Map(entries) = expr else {
192        return Err(Error::Eval(
193            "openai codec content part must be a map".to_owned(),
194        ));
195    };
196    match symbol_field(entries, "type")?.as_str() {
197        "text" => string_field(entries, "text"),
198        other => Err(Error::Eval(format!(
199            "openai codec does not support content part type {other}"
200        ))),
201    }
202}
203
204fn response_usage(entries: &[(Expr, Expr)]) -> Result<Value> {
205    let Some(usage) = entries.iter().find_map(|(field, value)| match field {
206        Expr::Symbol(symbol) if symbol.name.as_ref() == "usage" => Some(value),
207        _ => None,
208    }) else {
209        return Ok(Value::Null);
210    };
211    let Expr::Map(fields) = usage else {
212        return Err(Error::Eval(
213            "openai codec usage field must be a map".to_owned(),
214        ));
215    };
216    let prompt = optional_u64_field(fields, "input-tokens")?;
217    let completion = optional_u64_field(fields, "output-tokens")?;
218    let total = optional_u64_field(fields, "total-tokens")?.or_else(|| {
219        prompt
220            .zip(completion)
221            .map(|(left, right)| left.saturating_add(right))
222    });
223    Ok(json!({
224        "prompt_tokens": prompt.unwrap_or(0),
225        "completion_tokens": completion.unwrap_or(0),
226        "total_tokens": total.unwrap_or(0),
227    }))
228}