sim_lib_openai_server/routes/
threads.rs1use serde_json::{Value, json};
2
3use crate::{
4 clock::{SystemWallClock, WallClock},
5 ids::GatewayIdGenerator,
6 objects::{GatewayRequest, GatewayResponse},
7 server::GatewayRouteState,
8 storage::{GatewayStateStore, GatewayThread, GatewayThreadMessage},
9};
10
11use super::errors::OpenAiRouteError;
12
13pub const THREADS_PATH: &str = "/v1/threads";
15pub const THREAD_RETRIEVAL_PREFIX: &str = "/v1/threads/";
17pub const THREAD_RETRIEVAL_ROUTE: &str = "/v1/threads/{id}";
19pub const THREAD_MESSAGES_ROUTE: &str = "/v1/threads/{id}/messages";
21
22type RouteResult<T> = std::result::Result<T, OpenAiRouteError>;
23
24pub fn handle_threads(request: &GatewayRequest, state: &GatewayRouteState) -> GatewayResponse {
26 let mut clock = SystemWallClock;
27 let seed = clock.now_ms().unwrap_or(1);
28 let mut ids = GatewayIdGenerator::deterministic("thread", seed);
29 match state.store().lock() {
30 Ok(mut store) => create_thread(&mut *store, &mut ids, &mut clock, request)
31 .unwrap_or_else(OpenAiRouteError::into_response),
32 Err(err) => OpenAiRouteError::internal_message(format!("gateway store lock failed: {err}"))
33 .into_response(),
34 }
35}
36
37pub fn handle_thread_post(request: &GatewayRequest, state: &GatewayRouteState) -> GatewayResponse {
39 let Some(thread_id) = message_thread_id_from_path(request.path()) else {
40 return OpenAiRouteError::not_found_kind("thread", request.path()).into_response();
41 };
42 let mut clock = SystemWallClock;
43 let seed = clock.now_ms().unwrap_or(1);
44 let mut ids = GatewayIdGenerator::deterministic("msg", seed);
45 match state.store().lock() {
46 Ok(mut store) => append_message(&mut *store, &mut ids, &mut clock, thread_id, request)
47 .unwrap_or_else(OpenAiRouteError::into_response),
48 Err(err) => OpenAiRouteError::internal_message(format!("gateway store lock failed: {err}"))
49 .into_response(),
50 }
51}
52
53pub fn handle_thread_get(request: &GatewayRequest, state: &GatewayRouteState) -> GatewayResponse {
56 match thread_path(request.path()) {
57 Some(ThreadPath::Thread(thread_id)) => match state.store().lock() {
58 Ok(store) => retrieve_thread(&*store, thread_id),
59 Err(err) => {
60 OpenAiRouteError::internal_message(format!("gateway store lock failed: {err}"))
61 .into_response()
62 }
63 },
64 Some(ThreadPath::Messages(thread_id)) => match state.store().lock() {
65 Ok(store) => list_messages(&*store, thread_id),
66 Err(err) => {
67 OpenAiRouteError::internal_message(format!("gateway store lock failed: {err}"))
68 .into_response()
69 }
70 },
71 None => OpenAiRouteError::not_found_kind("thread", request.path()).into_response(),
72 }
73}
74
75pub fn retrieve_thread<S>(store: &S, thread_id: &str) -> GatewayResponse
77where
78 S: GatewayStateStore,
79{
80 store
81 .thread(thread_id)
82 .map(|thread| GatewayResponse::json_value(200, thread_json(&thread)))
83 .unwrap_or_else(|| OpenAiRouteError::not_found_kind("thread", thread_id).into_response())
84}
85
86pub fn list_messages<S>(store: &S, thread_id: &str) -> GatewayResponse
89where
90 S: GatewayStateStore,
91{
92 if store.thread(thread_id).is_none() {
93 return OpenAiRouteError::not_found_kind("thread", thread_id).into_response();
94 }
95 let data = store
96 .thread_messages(thread_id)
97 .iter()
98 .map(message_json)
99 .collect::<Vec<_>>();
100 GatewayResponse::json_value(200, json!({ "object": "list", "data": data }))
101}
102
103fn create_thread<S, C>(
104 store: &mut S,
105 ids: &mut GatewayIdGenerator,
106 clock: &mut C,
107 request: &GatewayRequest,
108) -> RouteResult<GatewayResponse>
109where
110 S: GatewayStateStore,
111 C: WallClock,
112{
113 let object = request_object(request.body())?;
114 let thread = GatewayThread::new(
115 ids.next_id().map_err(OpenAiRouteError::internal)?,
116 clock.now_ms().map_err(OpenAiRouteError::internal)?,
117 metadata(object.get("metadata"))?,
118 );
119 store
120 .put_thread(thread.clone())
121 .map_err(OpenAiRouteError::internal)?;
122 Ok(GatewayResponse::json_value(200, thread_json(&thread)))
123}
124
125fn append_message<S, C>(
126 store: &mut S,
127 ids: &mut GatewayIdGenerator,
128 clock: &mut C,
129 thread_id: &str,
130 request: &GatewayRequest,
131) -> RouteResult<GatewayResponse>
132where
133 S: GatewayStateStore,
134 C: WallClock,
135{
136 if store.thread(thread_id).is_none() {
137 return Err(OpenAiRouteError::not_found_kind("thread", thread_id));
138 }
139 let object = request_object(request.body())?;
140 let role = required_string(&object, "role")?.to_owned();
141 let content = required_string(&object, "content")?.to_owned();
142 let message = GatewayThreadMessage::new(
143 ids.next_id().map_err(OpenAiRouteError::internal)?,
144 thread_id,
145 role,
146 content,
147 clock.now_ms().map_err(OpenAiRouteError::internal)?,
148 );
149 store
150 .put_thread_message(message.clone())
151 .map_err(OpenAiRouteError::internal)?;
152 Ok(GatewayResponse::json_value(200, message_json(&message)))
153}
154
155use crate::routes::request_json::{request_object_or_empty as request_object, required_string};
156
157fn metadata(value: Option<&Value>) -> RouteResult<Vec<(String, String)>> {
158 let Some(value) = value else {
159 return Ok(Vec::new());
160 };
161 let object = value.as_object().ok_or_else(|| {
162 OpenAiRouteError::bad_request(
163 "metadata must be an object",
164 Some("metadata"),
165 "invalid_metadata",
166 )
167 })?;
168 object
169 .iter()
170 .map(|(key, value)| {
171 value
172 .as_str()
173 .map(|value| (key.clone(), value.to_owned()))
174 .ok_or_else(|| {
175 OpenAiRouteError::bad_request(
176 "metadata values must be strings",
177 Some("metadata"),
178 "invalid_metadata",
179 )
180 })
181 })
182 .collect()
183}
184
185enum ThreadPath<'a> {
186 Thread(&'a str),
187 Messages(&'a str),
188}
189
190fn thread_path(path: &str) -> Option<ThreadPath<'_>> {
191 if let Some(thread_id) = message_thread_id_from_path(path) {
192 return Some(ThreadPath::Messages(thread_id));
193 }
194 super::path::id_from_path(path, THREAD_RETRIEVAL_PREFIX).map(ThreadPath::Thread)
195}
196
197fn message_thread_id_from_path(path: &str) -> Option<&str> {
198 super::path::id_from_path_with_suffix(path, THREAD_RETRIEVAL_PREFIX, "/messages")
199}
200
201fn thread_json(thread: &GatewayThread) -> Value {
202 json!({
203 "id": thread.id(),
204 "object": "thread",
205 "created_at": thread.created_at_ms(),
206 "metadata": metadata_json(thread.metadata()),
207 })
208}
209
210fn message_json(message: &GatewayThreadMessage) -> Value {
211 json!({
212 "id": message.id(),
213 "object": "thread.message",
214 "thread_id": message.thread_id(),
215 "role": message.role(),
216 "content": message.content(),
217 "created_at": message.created_at_ms(),
218 })
219}
220
221fn metadata_json(metadata: &[(String, String)]) -> Value {
222 Value::Object(
223 metadata
224 .iter()
225 .map(|(key, value)| (key.clone(), Value::String(value.clone())))
226 .collect(),
227 )
228}