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