#![allow(missing_docs)]
use half::f16;
use serde::{Deserialize, Serialize};
use serde_json::Value;
use super::{RequestMetadata, TimingInfo};
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct SparseVector {
pub indices: Vec<u32>,
pub values: Vec<f32>,
}
impl SparseVector {
pub fn to_map(&self) -> std::collections::HashMap<u32, f32> {
self.indices
.iter()
.copied()
.zip(self.values.iter().copied())
.collect()
}
pub fn len(&self) -> usize {
self.indices.len().min(self.values.len())
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
}
#[derive(Debug, Clone, PartialEq)]
pub enum Multivector {
F16(Vec<Vec<f16>>),
F32(Vec<Vec<f32>>),
}
impl Multivector {
pub fn len(&self) -> usize {
match self {
Self::F16(rows) => rows.len(),
Self::F32(rows) => rows.len(),
}
}
pub fn is_empty(&self) -> bool {
self.len() == 0
}
pub fn dims(&self) -> usize {
match self {
Self::F16(rows) => rows.first().map_or(0, Vec::len),
Self::F32(rows) => rows.first().map_or(0, Vec::len),
}
}
pub fn to_f32(&self) -> Vec<Vec<f32>> {
match self {
Self::F16(rows) => rows
.iter()
.map(|row| row.iter().map(|value| value.to_f32()).collect())
.collect(),
Self::F32(rows) => rows.clone(),
}
}
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct EncodeResult {
pub model: Option<String>,
pub id: Option<String>,
pub dense: Option<Vec<f32>>,
pub sparse: Option<SparseVector>,
pub multivector: Option<Multivector>,
pub timing: Option<TimingInfo>,
pub request: Option<RequestMetadata>,
}
impl EncodeResult {
pub fn require_dense(&self) -> crate::error::Result<&[f32]> {
self.dense
.as_deref()
.ok_or_else(|| crate::error::Error::decode("encode result has no dense embedding"))
}
pub fn sparse_map(&self) -> std::collections::HashMap<u32, f32> {
self.sparse
.as_ref()
.map(SparseVector::to_map)
.unwrap_or_default()
}
pub fn multivector_f32(&self) -> Vec<Vec<f32>> {
self.multivector
.as_ref()
.map(Multivector::to_f32)
.unwrap_or_default()
}
}
#[derive(Debug, Clone, PartialEq, Serialize, Deserialize)]
pub struct ScoreEntry {
pub item_id: String,
pub score: f64,
pub rank: u32,
}
#[derive(Debug, Clone, Default, PartialEq, Eq, Serialize, Deserialize)]
pub struct ScoreUsage {
pub input_tokens: u64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub images: Option<u64>,
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct ScoreResult {
pub model: String,
pub query_id: Option<String>,
pub scores: Vec<ScoreEntry>,
pub usage: Option<ScoreUsage>,
pub request: Option<RequestMetadata>,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct Entity {
pub text: String,
pub label: String,
pub score: f64,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub start: Option<i64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub end: Option<i64>,
#[serde(default, skip_serializing_if = "Option::is_none")]
pub bbox: Option<Vec<i64>>,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct Relation {
pub head: String,
pub tail: String,
pub relation: String,
pub score: f64,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct Classification {
pub label: String,
pub score: f64,
}
#[derive(Debug, Clone, Default, PartialEq, Serialize, Deserialize)]
pub struct DetectedObject {
pub label: String,
pub score: f64,
pub bbox: Vec<i64>,
}
#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)]
pub struct ExtractItemError {
pub code: String,
pub message: String,
}
#[derive(Debug, Clone, Default, PartialEq)]
pub struct ExtractResult {
pub model: Option<String>,
pub id: Option<String>,
pub entities: Vec<Entity>,
pub relations: Vec<Relation>,
pub classifications: Vec<Classification>,
pub objects: Vec<DetectedObject>,
pub data: Option<Value>,
pub error: Option<ExtractItemError>,
pub request: Option<RequestMetadata>,
}
#[cfg(test)]
mod tests {
#![allow(clippy::float_cmp)]
use super::*;
#[test]
fn sparse_converts_to_a_term_weight_map() {
let sparse = SparseVector {
indices: vec![3, 9],
values: vec![0.5, 0.25],
};
let map = sparse.to_map();
assert_eq!(map.len(), 2);
assert_eq!(map[&9], 0.25);
assert_eq!(sparse.len(), 2);
assert!(SparseVector::default().is_empty());
}
#[test]
fn encode_result_accessors_are_lenient_except_where_they_promise_not_to_be() {
let empty = EncodeResult::default();
assert!(empty.require_dense().is_err());
assert!(empty.sparse_map().is_empty());
assert!(empty.multivector_f32().is_empty());
let filled = EncodeResult {
dense: Some(vec![0.5, 0.25]),
sparse: Some(SparseVector {
indices: vec![4],
values: vec![1.0],
}),
..EncodeResult::default()
};
assert_eq!(filled.require_dense().unwrap(), &[0.5, 0.25]);
assert_eq!(filled.sparse_map()[&4], 1.0);
}
#[test]
fn multivector_reports_its_shape_and_widens_on_request() {
let mv = Multivector::F16(vec![
vec![f16::from_f32(1.0), f16::from_f32(0.5)],
vec![f16::from_f32(0.0), f16::from_f32(-1.0)],
]);
assert_eq!(mv.len(), 2);
assert_eq!(mv.dims(), 2);
assert_eq!(mv.to_f32(), vec![vec![1.0, 0.5], vec![0.0, -1.0]]);
assert!(Multivector::F32(Vec::new()).is_empty());
}
}