1use std::collections::HashMap;
7
8use openkind_core::{Answer, Question, State, SystemRequest};
9use openkind_engine::{dispatch, EngineRegistry};
10use openkind_proto::openkind as pb;
11use tonic::{Request, Response, Status};
12
13use crate::middleware::AuthConfig;
14use crate::AppState;
15
16pub struct SystemOneService {
18 pub state: AppState,
20 pub auth: AuthConfig,
22}
23
24impl SystemOneService {
25 pub fn new(registry: EngineRegistry) -> Self {
27 Self::with_auth(registry, AuthConfig::default())
28 }
29
30 pub fn with_auth(registry: EngineRegistry, auth: AuthConfig) -> Self {
32 Self {
33 state: AppState::new(registry),
34 auth,
35 }
36 }
37}
38
39type RpcResult<T> = Result<Response<T>, Status>;
40
41#[tonic::async_trait]
42impl pb::system_one_server::SystemOne for SystemOneService {
43 async fn evaluate(
44 &self,
45 request: Request<pb::SystemOneRequest>,
46 ) -> RpcResult<pb::SystemOneResponse> {
47 let req_id = match request
48 .metadata()
49 .get("x-typesafe-request-id")
50 .and_then(|m| m.to_str().ok())
51 {
52 Some(id) if crate::middleware::is_safe_request_id(id) => id.to_string(),
53 _ => uuid::Uuid::new_v4().to_string(),
54 };
55
56 if self.auth.is_required() && !check_grpc_auth(request.metadata(), &self.auth) {
58 return Err(status_with_request_id(
59 Status::unauthenticated("missing or invalid API key"),
60 &req_id,
61 ));
62 }
63
64 let pb_req = request.into_inner();
65
66 let state = pb_state_to_core(pb_req.state.as_ref())
68 .map_err(|status| status_with_request_id(status, &req_id))?;
69 let questions = pb_questions_to_core(pb_req.questions)
70 .map_err(|status| status_with_request_id(status, &req_id))?;
71 let req = SystemRequest {
72 state,
73 model: pb_req.model,
74 questions,
75 };
76
77 let resp = dispatch(req, &self.state.registry)
78 .await
79 .map_err(|error| status_with_request_id(status_from_engine(error), &req_id))?;
80
81 let mut response = Response::new(core_to_pb_response(resp));
82 response.metadata_mut().append(
86 "x-typesafe-request-id",
87 req_id
88 .parse()
89 .unwrap_or_else(|_| tonic::metadata::MetadataValue::from_static("invalid")),
90 );
91 Ok(response)
92 }
93}
94
95fn status_with_request_id(mut status: Status, request_id: &str) -> Status {
96 if let Ok(value) = request_id.parse() {
99 status.metadata_mut().insert("x-typesafe-request-id", value);
100 }
101 status
102}
103
104fn check_grpc_auth(metadata: &tonic::metadata::MetadataMap, auth: &AuthConfig) -> bool {
105 if !auth.is_required() {
106 return true;
107 }
108
109 let supplied = metadata
110 .get("authorization")
111 .and_then(|v| v.to_str().ok())
112 .and_then(|s| {
113 let (scheme, token) = s.split_once(' ')?;
114 scheme.eq_ignore_ascii_case("Bearer").then_some(token)
115 })
116 .or_else(|| metadata.get("x-api-key").and_then(|v| v.to_str().ok()));
117
118 match supplied {
119 Some(token) => auth.token_matches(token),
120 None => false,
121 }
122}
123
124fn pb_state_to_core(pb: Option<&pb::State>) -> Result<State, Status> {
127 let pb = pb.ok_or_else(|| Status::invalid_argument("state is required"))?;
128 match &pb.value {
129 Some(pb::state::Value::Text(s)) => Ok(State::Text(s.clone())),
130 Some(pb::state::Value::Structured(s)) => {
131 let v: serde_json::Value = serde_json::from_slice(&s.json)
134 .map_err(|e| Status::invalid_argument(e.to_string()))?;
135 json_value_to_state(v)
136 }
137 None => Err(Status::invalid_argument("state must be set")),
138 }
139}
140
141fn json_value_to_state(v: serde_json::Value) -> Result<State, Status> {
142 match v {
143 serde_json::Value::String(s) => Ok(State::Text(s)),
144 serde_json::Value::Object(m) => Ok(State::Object(m)),
145 serde_json::Value::Array(a) => Ok(State::Array(a)),
146 _ => Err(Status::invalid_argument(
147 "structured state must be a JSON object, array, or string",
148 )),
149 }
150}
151
152fn pb_questions_to_core(
153 pb: HashMap<String, pb::Question>,
154) -> Result<HashMap<String, Question, openkind_core::WireHashState>, Status> {
155 if pb.len() > openkind_core::MAX_QUESTIONS_PER_REQUEST {
156 return Err(Status::invalid_argument(format!(
157 "request exceeds maximum question count limit (got {}, max {})",
158 pb.len(),
159 openkind_core::MAX_QUESTIONS_PER_REQUEST
160 )));
161 }
162 let cap = pb.len().min(openkind_core::MAX_QUESTIONS_PER_REQUEST);
163 let mut out: HashMap<String, Question, openkind_core::WireHashState> =
164 HashMap::with_capacity_and_hasher(cap, Default::default());
165 for (id, q) in pb {
166 let kind = q
167 .kind
168 .ok_or_else(|| Status::invalid_argument("question has no kind"))?;
169 let core = match kind {
170 pb::question::Kind::Noul(n) => {
171 let instr = parse_json(&n.instructions_json)?;
172 let criteria = n.criteria.map(|c| openkind_core::NoulCriteria {
173 r#true: c.is_true,
174 r#false: c.is_false,
175 });
176 Question::Noul(openkind_core::NoulQuestion {
177 instructions: instr,
178 criteria,
179 })
180 }
181 pb::question::Kind::Choice(c) => {
182 if c.criteria.len() > openkind_core::MAX_CRITERIA_OPTIONS {
183 return Err(Status::invalid_argument(format!(
184 "choice question `{id}` exceeds maximum criteria options limit (got {}, max {})",
185 c.criteria.len(),
186 openkind_core::MAX_CRITERIA_OPTIONS
187 )));
188 }
189 let instr = parse_json(&c.instructions_json)?;
190 let criteria = c
191 .criteria
192 .into_iter()
193 .map(|(k, v)| (k, if v.is_empty() { None } else { Some(v) }))
194 .collect();
195 Question::Choice(openkind_core::ChoiceQuestion {
196 instructions: instr,
197 criteria,
198 })
199 }
200 pb::question::Kind::Score(s) => {
201 if s.criteria.len() > openkind_core::MAX_CRITERIA_OPTIONS {
202 return Err(Status::invalid_argument(format!(
203 "score question `{id}` exceeds maximum criteria options limit (got {}, max {})",
204 s.criteria.len(),
205 openkind_core::MAX_CRITERIA_OPTIONS
206 )));
207 }
208 let instr = parse_json(&s.instructions_json)?;
209 Question::Score(openkind_core::ScoreQuestion {
210 instructions: instr,
211 criteria: s.criteria,
212 })
213 }
214 };
215 out.insert(id, core);
216 }
217 Ok(out)
218}
219
220fn parse_json(bytes: &[u8]) -> Result<serde_json::Value, Status> {
221 if bytes.is_empty() {
222 return Ok(serde_json::Value::Null);
223 }
224 serde_json::from_slice(bytes).map_err(|e| Status::invalid_argument(e.to_string()))
225}
226
227fn core_to_pb_response(resp: openkind_core::SystemResponse) -> pb::SystemOneResponse {
228 let answers = resp
229 .answers
230 .into_iter()
231 .map(|(id, ans)| {
232 let pb_ans = match ans {
233 Answer::Noul(n) => pb::Answer {
234 kind: Some(pb::answer::Kind::Noul(pb::NoulAnswer { noul: n.noul })),
235 },
236 Answer::Choice(c) => pb::Answer {
237 kind: Some(pb::answer::Kind::Choice(pb::ChoiceAnswer {
238 choice: c.choice,
239 probabilities: c.probabilities,
240 confidence: c.confidence,
241 })),
242 },
243 Answer::Score(s) => pb::Answer {
244 kind: Some(pb::answer::Kind::Score(pb::ScoreAnswer {
245 score: s.score,
246 legend: s.legend,
247 probabilities: s.probabilities,
248 confidence: s.confidence,
249 })),
250 },
251 };
252 (id, pb_ans)
253 })
254 .collect();
255 pb::SystemOneResponse {
256 model: resp.model,
257 answers,
258 usage: Some(pb::Usage {
259 input_tokens: resp.usage.input_tokens,
260 output_tokens: resp.usage.output_tokens,
261 }),
262 }
263}
264
265fn status_from_engine(e: openkind_engine::EngineError) -> Status {
266 use openkind_engine::EngineError::*;
267 match e {
268 Invalid(_) => Status::invalid_argument(e.to_string()),
269 UnknownModel(_) => Status::not_found(e.to_string()),
270 Unsupported { .. } => Status::invalid_argument(e.to_string()),
271 Overloaded { .. } => Status::unavailable(e.to_string()),
272 DeadlineExceeded { .. } => Status::deadline_exceeded(e.to_string()),
273 Backend { .. } => Status::internal(e.to_string()),
274 }
275}
276
277pub fn server(
279 registry: EngineRegistry,
280) -> pb::system_one_server::SystemOneServer<SystemOneService> {
281 server_with_auth(registry, AuthConfig::default())
282}
283
284pub fn server_with_auth(
286 registry: EngineRegistry,
287 auth: AuthConfig,
288) -> pb::system_one_server::SystemOneServer<SystemOneService> {
289 pb::system_one_server::SystemOneServer::new(SystemOneService::with_auth(registry, auth))
290 .max_decoding_message_size(16 * 1024 * 1024)
291 .max_encoding_message_size(16 * 1024 * 1024)
292}
293
294pub fn service(
296 registry: EngineRegistry,
297) -> pb::system_one_server::SystemOneServer<SystemOneService> {
298 server(registry)
299}
300
301pub fn service_with_auth(
303 registry: EngineRegistry,
304 auth: AuthConfig,
305) -> pb::system_one_server::SystemOneServer<SystemOneService> {
306 server_with_auth(registry, auth)
307}