use async_trait::async_trait;
use openkind_core::{ModelInfo, SystemRequest, SystemResponse};
use crate::error::EngineResult;
#[async_trait]
pub trait DecisionEngine: Send + Sync {
fn backend_id(&self) -> &str;
fn model_metadata(&self) -> ModelInfo {
ModelInfo {
name: self.backend_id().to_string(),
description: format!("{} backend", self.backend_id()),
release_date: "1970-01-01".to_string(),
}
}
async fn evaluate(&self, req: SystemRequest) -> EngineResult<SystemResponse>;
fn estimate_input_tokens(&self, req: &SystemRequest) -> u32 {
let state_chars = match &req.state {
openkind_core::State::Text(s) => s.len(),
openkind_core::State::Object(m) => count_json_bytes(m),
openkind_core::State::Array(a) => count_json_bytes(a),
};
let question_chars =
req.questions
.values()
.map(|q| match q {
openkind_core::Question::Noul(n) => count_json_bytes(&n.instructions)
.saturating_add(n.criteria.as_ref().map(count_json_bytes).unwrap_or(0)),
openkind_core::Question::Choice(c) => count_json_bytes(&c.instructions)
.saturating_add(count_json_bytes(&c.criteria)),
openkind_core::Question::Score(s) => count_json_bytes(&s.instructions)
.saturating_add(count_json_bytes(&s.criteria)),
})
.fold(0usize, usize::saturating_add);
let total_chars = state_chars.saturating_add(question_chars);
if total_chars == 0 {
0
} else {
u32::try_from(total_chars.div_ceil(4)).unwrap_or(u32::MAX)
}
}
}
struct ByteCounter(usize);
impl std::io::Write for ByteCounter {
fn write(&mut self, buf: &[u8]) -> std::io::Result<usize> {
self.0 = self.0.saturating_add(buf.len());
Ok(buf.len())
}
fn flush(&mut self) -> std::io::Result<()> {
Ok(())
}
}
fn count_json_bytes<T: serde::Serialize + ?Sized>(val: &T) -> usize {
let mut counter = ByteCounter(0);
serde_json::to_writer(&mut counter, val)
.map(|()| counter.0)
.unwrap_or(0)
}