use std::collections::HashMap;
use openkind_core::{Answer, Question, State, SystemRequest};
use openkind_engine::{dispatch, EngineRegistry};
use openkind_proto::openkind as pb;
use tonic::{Request, Response, Status};
use crate::middleware::{AuthConfig, RequestLimits};
use crate::AppState;
pub struct SystemOneService {
pub state: AppState,
pub auth: AuthConfig,
pub limits: RequestLimits,
}
impl SystemOneService {
pub fn new(registry: EngineRegistry) -> Self {
Self::with_auth(registry, AuthConfig::default())
}
pub fn with_auth(registry: EngineRegistry, auth: AuthConfig) -> Self {
Self::with_auth_and_limits(registry, auth, RequestLimits::default())
}
pub fn with_auth_and_limits(
registry: EngineRegistry,
auth: AuthConfig,
limits: RequestLimits,
) -> Self {
Self {
state: AppState::new(registry),
auth,
limits,
}
}
}
type RpcResult<T> = Result<Response<T>, Status>;
#[tonic::async_trait]
impl pb::system_one_server::SystemOne for SystemOneService {
async fn evaluate(
&self,
request: Request<pb::SystemOneRequest>,
) -> RpcResult<pb::SystemOneResponse> {
let req_id = match request
.metadata()
.get("x-typesafe-request-id")
.and_then(|m| m.to_str().ok())
{
Some(id) if crate::middleware::is_safe_request_id(id) => id.to_string(),
_ => uuid::Uuid::new_v4().to_string(),
};
if self.auth.is_required() && !check_grpc_auth(request.metadata(), &self.auth) {
metrics::counter!("openkind_auth_failures_total", "transport" => "grpc").increment(1);
if let Some(peer) = request.remote_addr() {
if let Err(ms) = self.limits.failed_auth.check(peer.ip()) {
return Err(status_with_request_id(
status_with_retry(
Status::resource_exhausted("authentication rate limited"),
ms,
),
&req_id,
));
}
}
return Err(status_with_request_id(
Status::unauthenticated("missing or invalid API key"),
&req_id,
));
}
if let Some(peer) = request.remote_addr() {
if let Err(ms) = self.limits.evaluation.check(peer.ip()) {
return Err(status_with_request_id(
status_with_retry(Status::resource_exhausted("rate limited"), ms),
&req_id,
));
}
}
let pb_req = request.into_inner();
let state = pb_state_to_core(pb_req.state.as_ref())
.map_err(|status| status_with_request_id(status, &req_id))?;
let questions = pb_questions_to_core(pb_req.questions)
.map_err(|status| status_with_request_id(status, &req_id))?;
let req = SystemRequest {
state,
model: pb_req.model,
questions,
};
let resp = dispatch(req, &self.state.registry)
.await
.map_err(|error| status_with_request_id(status_from_engine(error), &req_id))?;
let mut response = Response::new(core_to_pb_response(resp));
response.metadata_mut().append(
"x-typesafe-request-id",
req_id
.parse()
.unwrap_or_else(|_| tonic::metadata::MetadataValue::from_static("invalid")),
);
Ok(response)
}
}
fn status_with_retry(mut status: Status, ms: u64) -> Status {
if let Ok(value) = ms.to_string().parse() {
status.metadata_mut().insert("retry-after-ms", value);
}
if let Ok(value) = ms.div_ceil(1000).to_string().parse() {
status.metadata_mut().insert("retry-after", value);
}
status
}
fn status_with_request_id(mut status: Status, request_id: &str) -> Status {
if let Ok(value) = request_id.parse() {
status.metadata_mut().insert("x-typesafe-request-id", value);
}
status
}
fn check_grpc_auth(metadata: &tonic::metadata::MetadataMap, auth: &AuthConfig) -> bool {
if !auth.is_required() {
return true;
}
let supplied = metadata
.get("authorization")
.and_then(|v| v.to_str().ok())
.and_then(|s| {
let (scheme, token) = s.split_once(' ')?;
scheme.eq_ignore_ascii_case("Bearer").then_some(token)
})
.or_else(|| metadata.get("x-api-key").and_then(|v| v.to_str().ok()));
match supplied {
Some(token) => auth.token_matches(token),
None => false,
}
}
fn pb_state_to_core(pb: Option<&pb::State>) -> Result<State, Status> {
let pb = pb.ok_or_else(|| Status::invalid_argument("state is required"))?;
match &pb.value {
Some(pb::state::Value::Text(s)) => Ok(State::Text(s.clone())),
Some(pb::state::Value::Structured(s)) => {
let v: serde_json::Value = serde_json::from_slice(&s.json)
.map_err(|e| Status::invalid_argument(e.to_string()))?;
json_value_to_state(v)
}
None => Err(Status::invalid_argument("state must be set")),
}
}
fn json_value_to_state(v: serde_json::Value) -> Result<State, Status> {
match v {
serde_json::Value::String(s) => Ok(State::Text(s)),
serde_json::Value::Object(m) => Ok(State::Object(m)),
serde_json::Value::Array(a) => Ok(State::Array(a)),
_ => Err(Status::invalid_argument(
"structured state must be a JSON object, array, or string",
)),
}
}
fn pb_questions_to_core(
pb: HashMap<String, pb::Question>,
) -> Result<HashMap<String, Question, openkind_core::WireHashState>, Status> {
if pb.len() > openkind_core::MAX_QUESTIONS_PER_REQUEST {
return Err(Status::invalid_argument(format!(
"request exceeds maximum question count limit (got {}, max {})",
pb.len(),
openkind_core::MAX_QUESTIONS_PER_REQUEST
)));
}
let cap = pb.len().min(openkind_core::MAX_QUESTIONS_PER_REQUEST);
let mut out: HashMap<String, Question, openkind_core::WireHashState> =
HashMap::with_capacity_and_hasher(cap, Default::default());
for (id, q) in pb {
let kind = q
.kind
.ok_or_else(|| Status::invalid_argument("question has no kind"))?;
let core = match kind {
pb::question::Kind::Noul(n) => {
let instr = parse_json(&n.instructions_json)?;
let criteria = n.criteria.map(|c| openkind_core::NoulCriteria {
r#true: c.is_true,
r#false: c.is_false,
});
Question::Noul(openkind_core::NoulQuestion {
instructions: instr,
criteria,
})
}
pb::question::Kind::Choice(c) => {
if c.criteria.len() > openkind_core::MAX_CRITERIA_OPTIONS {
return Err(Status::invalid_argument(format!(
"choice question `{id}` exceeds maximum criteria options limit (got {}, max {})",
c.criteria.len(),
openkind_core::MAX_CRITERIA_OPTIONS
)));
}
let instr = parse_json(&c.instructions_json)?;
let criteria = c
.criteria
.into_iter()
.map(|(k, v)| (k, if v.is_empty() { None } else { Some(v) }))
.collect();
Question::Choice(openkind_core::ChoiceQuestion {
instructions: instr,
criteria,
})
}
pb::question::Kind::Score(s) => {
if s.criteria.len() > openkind_core::MAX_CRITERIA_OPTIONS {
return Err(Status::invalid_argument(format!(
"score question `{id}` exceeds maximum criteria options limit (got {}, max {})",
s.criteria.len(),
openkind_core::MAX_CRITERIA_OPTIONS
)));
}
let instr = parse_json(&s.instructions_json)?;
Question::Score(openkind_core::ScoreQuestion {
instructions: instr,
criteria: s.criteria,
})
}
};
out.insert(id, core);
}
Ok(out)
}
fn parse_json(bytes: &[u8]) -> Result<serde_json::Value, Status> {
if bytes.is_empty() {
return Ok(serde_json::Value::Null);
}
serde_json::from_slice(bytes).map_err(|e| Status::invalid_argument(e.to_string()))
}
fn core_to_pb_response(resp: openkind_core::SystemResponse) -> pb::SystemOneResponse {
let answers = resp
.answers
.into_iter()
.map(|(id, ans)| {
let pb_ans = match ans {
Answer::Noul(n) => pb::Answer {
kind: Some(pb::answer::Kind::Noul(pb::NoulAnswer { noul: n.noul })),
},
Answer::Choice(c) => pb::Answer {
kind: Some(pb::answer::Kind::Choice(pb::ChoiceAnswer {
choice: c.choice,
probabilities: c.probabilities,
confidence: c.confidence,
})),
},
Answer::Score(s) => pb::Answer {
kind: Some(pb::answer::Kind::Score(pb::ScoreAnswer {
score: s.score,
legend: s.legend,
probabilities: s.probabilities,
confidence: s.confidence,
})),
},
};
(id, pb_ans)
})
.collect();
pb::SystemOneResponse {
model: resp.model,
answers,
usage: Some(pb::Usage {
input_tokens: resp.usage.input_tokens,
output_tokens: resp.usage.output_tokens,
}),
}
}
fn status_from_engine(e: openkind_engine::EngineError) -> Status {
use openkind_engine::EngineError::*;
match e {
Invalid(_) => Status::invalid_argument(e.to_string()),
UnknownModel(_) => Status::not_found(e.to_string()),
Unsupported { .. } => Status::invalid_argument(e.to_string()),
Overloaded { retry_after_ms, .. } => {
status_with_retry(Status::unavailable(e.to_string()), retry_after_ms)
}
DeadlineExceeded { .. } => Status::deadline_exceeded(e.to_string()),
Backend { .. } | BackendValidation { .. } => Status::internal(e.to_string()),
}
}
pub fn server(
registry: EngineRegistry,
) -> pb::system_one_server::SystemOneServer<SystemOneService> {
server_with_auth(registry, AuthConfig::default())
}
pub fn server_with_auth(
registry: EngineRegistry,
auth: AuthConfig,
) -> pb::system_one_server::SystemOneServer<SystemOneService> {
pb::system_one_server::SystemOneServer::new(SystemOneService::with_auth(registry, auth))
.max_decoding_message_size(16 * 1024 * 1024)
.max_encoding_message_size(16 * 1024 * 1024)
}
pub fn service(
registry: EngineRegistry,
) -> pb::system_one_server::SystemOneServer<SystemOneService> {
server(registry)
}
pub fn service_with_auth(
registry: EngineRegistry,
auth: AuthConfig,
) -> pb::system_one_server::SystemOneServer<SystemOneService> {
server_with_auth(registry, auth)
}
pub fn service_with_auth_and_limits(
registry: EngineRegistry,
auth: AuthConfig,
limits: RequestLimits,
) -> pb::system_one_server::SystemOneServer<SystemOneService> {
pb::system_one_server::SystemOneServer::new(SystemOneService::with_auth_and_limits(
registry, auth, limits,
))
.max_decoding_message_size(16 * 1024 * 1024)
.max_encoding_message_size(16 * 1024 * 1024)
}