use super::jev;
use super::jev_cache::Caches;
use super::pool;
use crate::core::config::DecisionConfig;
use crate::core::decision::{self, Asked, Batch, DecisionSpec, Unit};
#[derive(Clone)]
pub struct Judged {
pub probability: f64,
pub model: String,
pub cached: bool,
}
pub type Answer = Result<Judged, String>;
pub struct Check<'a> {
pub id: &'a str,
pub spec: &'a DecisionSpec,
pub jev: &'a DecisionConfig,
pub units: &'a [Unit],
}
pub fn answers(caches: Result<&mut Caches, &String>, checks: &[Check]) -> Vec<Vec<Answer>> {
let sizes: Vec<usize> = checks.iter().map(|c| c.units.len()).collect();
if sizes.iter().all(|&n| n == 0) {
return checks.iter().map(|_| Vec::new()).collect();
}
let caches = match caches {
Ok(caches) => caches,
Err(e) => return sizes.iter().map(|&n| vec![Err(e.clone()); n]).collect(),
};
let mut groups: Vec<(&DecisionConfig, Vec<usize>)> = Vec::new();
for (i, check) in checks.iter().enumerate() {
match groups.iter_mut().find(|(config, _)| *config == check.jev) {
Some((_, members)) => members.push(i),
None => groups.push((check.jev, vec![i])),
}
}
let mut answers: Vec<Option<Vec<Answer>>> = checks.iter().map(|_| None).collect();
for (config, members) in groups {
let group: Vec<&Check> = members.iter().map(|&i| &checks[i]).collect();
for (&i, check_answers) in members.iter().zip(answer_group(caches, config, &group)) {
answers[i] = Some(check_answers);
}
}
answers
.into_iter()
.map(|a| a.expect("every check belongs to a group"))
.collect()
}
fn answer_group(
caches: &mut Caches,
config: &DecisionConfig,
checks: &[&Check],
) -> Vec<Vec<Answer>> {
let units: Vec<&[Unit]> = checks.iter().map(|c| c.units).collect();
let sizes: Vec<usize> = units.iter().map(|u| u.len()).collect();
let batches = decision::batch(&units);
let per_batch = match ask_batches(caches, config, checks, &batches) {
Ok(per_batch) => per_batch,
Err(e) => batches
.iter()
.map(|b| vec![Err(e.clone()); b.members.len()])
.collect(),
};
decision::spread(&batches, &sizes, per_batch)
}
fn ask_batches(
caches: &mut Caches,
config: &DecisionConfig,
checks: &[&Check],
batches: &[Batch],
) -> Result<Vec<Vec<Answer>>, String> {
if batches.is_empty() {
return Ok(Vec::new());
}
let client = connect(config)?;
let model = resolve_model(&client, config)?;
let asked: Vec<Asked> = batches
.iter()
.map(|batch| {
let members: Vec<(&str, &DecisionSpec)> = batch
.members
.iter()
.map(|&(check, _)| (checks[check].id, checks[check].spec))
.collect();
decision::ask(batch.unit, &members, &model)
})
.collect();
let mut results: Vec<Vec<Option<Answer>>> = batches
.iter()
.zip(&asked)
.map(|(batch, asked)| from_cache(caches, batch, asked, &model))
.collect();
let pending: Vec<usize> = (0..batches.len())
.filter(|&i| results[i].iter().any(Option::is_none))
.collect();
let responses = pool::map(&pending, config.concurrency, |&i| {
send(&client, config, &asked[i], &results[i])
});
for (&i, response) in pending.iter().zip(responses) {
record(caches, &asked[i], &mut results[i], response);
}
Ok(results
.into_iter()
.map(|r| {
r.into_iter()
.map(|a| a.expect("every question was answered or failed"))
.collect()
})
.collect())
}
pub(super) fn connect(config: &DecisionConfig) -> Result<jev::Client, String> {
match std::env::var(&config.api_key_env) {
Ok(key) if !key.is_empty() => Ok(jev::Client::new(&config.endpoint, key)),
_ => Err(format!(
"{} is not set; decision checks call Jev with it",
config.api_key_env
)),
}
}
pub(super) fn resolve_model(
client: &jev::Client,
config: &DecisionConfig,
) -> Result<String, String> {
let body = client.ask(&decision::probe_request(&config.model))?;
decision::parse_response(&body, &[]).map(|(model, _)| model)
}
fn from_cache(
caches: &mut Caches,
batch: &Batch,
asked: &Asked,
model: &str,
) -> Vec<Option<Answer>> {
if let Some(reason) = decision::too_large(batch.unit) {
return vec![Some(Err(reason)); asked.keys.len()];
}
asked
.keys
.iter()
.map(|key| {
caches.get(key).map(|probability| {
Ok(Judged {
probability,
model: model.to_string(),
cached: true,
})
})
})
.collect()
}
pub(super) fn send(
client: &jev::Client,
config: &DecisionConfig,
asked: &Asked,
cached: &[Option<Answer>],
) -> Result<(String, Vec<f64>), String> {
let unanswered: Vec<&(String, serde_json::Value)> = asked
.questions
.questions
.iter()
.zip(cached)
.filter(|(_, answer)| answer.is_none())
.map(|(question, _)| question)
.collect();
let ids: Vec<&str> = unanswered.iter().map(|(id, _)| id.as_str()).collect();
let request = decision::request(&config.model, &asked.questions.state, &unanswered);
client
.ask(&request)
.and_then(|body| decision::parse_response(&body, &ids))
}
pub(super) fn record(
caches: &mut Caches,
asked: &Asked,
results: &mut [Option<Answer>],
response: Result<(String, Vec<f64>), String>,
) {
let slots = results
.iter_mut()
.zip(&asked.questions.questions)
.filter(|(result, _)| result.is_none());
match response {
Ok((answered_by, probabilities)) => {
for ((slot, (id, question)), p) in slots.zip(probabilities) {
let key = decision::cache_key(&answered_by, &asked.questions.state, id, question);
caches.insert(key, p);
*slot = Some(Ok(Judged {
probability: p,
model: answered_by.clone(),
cached: false,
}));
}
}
Err(e) => {
for (slot, _) in slots {
*slot = Some(Err(e.clone()));
}
}
}
}