Skip to main content

sim_lib_openai_server/routes/
threads.rs

1use 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
13/// Route path for thread creation (`POST /v1/threads`).
14pub const THREADS_PATH: &str = "/v1/threads";
15/// Path prefix shared by thread retrieval and message routes (`/v1/threads/`).
16pub const THREAD_RETRIEVAL_PREFIX: &str = "/v1/threads/";
17/// Templated route for retrieving a single thread by id (`/v1/threads/{id}`).
18pub const THREAD_RETRIEVAL_ROUTE: &str = "/v1/threads/{id}";
19/// Templated route for a thread's message collection (`/v1/threads/{id}/messages`).
20pub const THREAD_MESSAGES_ROUTE: &str = "/v1/threads/{id}/messages";
21
22type RouteResult<T> = std::result::Result<T, OpenAiRouteError>;
23
24/// Handles `POST /v1/threads`, creating a new thread and returning its JSON object.
25pub 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
37/// Handles `POST /v1/threads/{id}/messages`, appending a message to the thread.
38pub 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
53/// Handles `GET /v1/threads/{id}` and `GET /v1/threads/{id}/messages`, dispatching
54/// to thread retrieval or message listing based on the request path.
55pub 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
75/// Returns the JSON object for a stored thread, or a not-found error response.
76pub 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
86/// Returns the thread's messages as an OpenAI `list` object, or a not-found error
87/// response when the thread does not exist.
88pub 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}