Skip to main content

clt_database/vector/
mod.rs

1use crate::types::AsValueRef;
2use crate::types::Value;
3use crate::types::ValueType;
4use crate::vdbe::Register;
5use crate::LimboError;
6use crate::Result;
7use crate::ValueRef;
8
9pub mod operations;
10pub mod vector_types;
11use vector_types::*;
12
13pub fn parse_vector<'a>(
14    value: &'a (impl AsValueRef + 'a),
15    type_hint: Option<VectorType>,
16) -> Result<Vector<'a>> {
17    let value = value.as_value_ref();
18    match value.value_type() {
19        ValueType::Text => operations::text::vector_from_text(
20            type_hint.unwrap_or(VectorType::Float32Dense),
21            value.to_text().expect("value must be text"),
22        ),
23        ValueType::Blob => {
24            let Some(blob) = value.to_blob() else {
25                return Err(LimboError::ConversionError(
26                    "Invalid vector value".to_string(),
27                ));
28            };
29            Vector::from_slice(blob)
30        }
31        _ => Err(LimboError::ConversionError(
32            "Invalid vector type".to_string(),
33        )),
34    }
35}
36
37pub fn vector32(args: &[Register]) -> Result<Value> {
38    if args.len() != 1 {
39        return Err(LimboError::ConversionError(
40            "vector32 requires exactly one argument".to_string(),
41        ));
42    }
43    let value = args[0].get_value();
44    let vector = parse_vector(value, Some(VectorType::Float32Dense))?;
45    let vector = operations::convert::vector_convert(vector, VectorType::Float32Dense)?;
46    Ok(operations::serialize::vector_serialize(vector))
47}
48
49pub fn vector32_sparse(args: &[Register]) -> Result<Value> {
50    if args.len() != 1 {
51        return Err(LimboError::ConversionError(
52            "vector32_sparse requires exactly one argument".to_string(),
53        ));
54    }
55    let value = args[0].get_value();
56    let vector = parse_vector(value, Some(VectorType::Float32Sparse))?;
57    let vector = operations::convert::vector_convert(vector, VectorType::Float32Sparse)?;
58    Ok(operations::serialize::vector_serialize(vector))
59}
60
61pub fn vector64(args: &[Register]) -> Result<Value> {
62    if args.len() != 1 {
63        return Err(LimboError::ConversionError(
64            "vector64 requires exactly one argument".to_string(),
65        ));
66    }
67    let value = args[0].get_value();
68    let vector = parse_vector(value, Some(VectorType::Float64Dense))?;
69    let vector = operations::convert::vector_convert(vector, VectorType::Float64Dense)?;
70    Ok(operations::serialize::vector_serialize(vector))
71}
72
73pub fn vector8(args: &[Register]) -> Result<Value> {
74    if args.len() != 1 {
75        return Err(LimboError::ConversionError(
76            "vector8 requires exactly one argument".to_string(),
77        ));
78    }
79    let value = args[0].get_value();
80    let vector = parse_vector(value, Some(VectorType::Float8))?;
81    let vector = operations::convert::vector_convert(vector, VectorType::Float8)?;
82    Ok(operations::serialize::vector_serialize(vector))
83}
84
85pub fn vector1bit(args: &[Register]) -> Result<Value> {
86    if args.len() != 1 {
87        return Err(LimboError::ConversionError(
88            "vector1bit requires exactly one argument".to_string(),
89        ));
90    }
91    let value = args[0].get_value();
92    let vector = parse_vector(value, Some(VectorType::Float1Bit))?;
93    let vector = operations::convert::vector_convert(vector, VectorType::Float1Bit)?;
94    Ok(operations::serialize::vector_serialize(vector))
95}
96
97pub fn vector_extract(args: &[Register]) -> Result<Value> {
98    if args.len() != 1 {
99        return Err(LimboError::ConversionError(
100            "vector_extract requires exactly one argument".to_string(),
101        ));
102    }
103
104    let value = args[0].get_value().as_value_ref();
105    let blob = match value {
106        ValueRef::Blob(b) => b,
107        _ => {
108            return Err(LimboError::ConversionError(
109                "Expected blob value".to_string(),
110            ))
111        }
112    };
113
114    if blob.is_empty() {
115        return Ok(Value::build_text("[]"));
116    }
117
118    let vector = Vector::from_slice(blob)?;
119    Ok(Value::build_text(operations::text::vector_to_text(&vector)))
120}
121
122pub fn vector_distance_cos(args: &[Register]) -> Result<Value> {
123    if args.len() != 2 {
124        return Err(LimboError::ConversionError(
125            "vector_distance_cos requires exactly two arguments".to_string(),
126        ));
127    }
128
129    let value_0 = args[0].get_value();
130    let value_1 = args[1].get_value();
131    let x = parse_vector(value_0, None)?;
132    let y = parse_vector(value_1, None)?;
133    let dist = operations::distance_cos::vector_distance_cos(&x, &y)?;
134    Ok(Value::from_f64(dist))
135}
136
137pub fn vector_distance_l2(args: &[Register]) -> Result<Value> {
138    if args.len() != 2 {
139        return Err(LimboError::ConversionError(
140            "distance_l2 requires exactly two arguments".to_string(),
141        ));
142    }
143
144    let value_0 = args[0].get_value();
145    let value_1 = args[1].get_value();
146    let x = parse_vector(value_0, None)?;
147    let y = parse_vector(value_1, None)?;
148    let dist = operations::distance_l2::vector_distance_l2(&x, &y)?;
149    Ok(Value::from_f64(dist))
150}
151
152pub fn vector_distance_jaccard(args: &[Register]) -> Result<Value> {
153    if args.len() != 2 {
154        return Err(LimboError::ConversionError(
155            "distance_jaccard requires exactly two arguments".to_string(),
156        ));
157    }
158
159    let value_0 = args[0].get_value();
160    let value_1 = args[1].get_value();
161    let x = parse_vector(value_0, None)?;
162    let y = parse_vector(value_1, None)?;
163    let dist = operations::jaccard::vector_distance_jaccard(&x, &y)?;
164    Ok(Value::from_f64(dist))
165}
166
167pub fn vector_distance_dot(args: &[Register]) -> Result<Value> {
168    if args.len() != 2 {
169        return Err(LimboError::ConversionError(
170            "distance_dot requires exactly two arguments".to_string(),
171        ));
172    }
173
174    let value_0 = args[0].get_value();
175    let value_1 = args[1].get_value();
176    let x = parse_vector(value_0, None)?;
177    let y = parse_vector(value_1, None)?;
178    let dist = operations::distance_dot::vector_distance_dot(&x, &y)?;
179    Ok(Value::from_f64(dist))
180}
181
182pub fn vector_concat(args: &[Register]) -> Result<Value> {
183    if args.len() != 2 {
184        return Err(LimboError::InvalidArgument(
185            "concat requires exactly two arguments".into(),
186        ));
187    }
188
189    let value_0 = args[0].get_value();
190    let value_1 = args[1].get_value();
191    let x = parse_vector(value_0, None)?;
192    let y = parse_vector(value_1, None)?;
193    let vector = operations::concat::vector_concat(&x, &y)?;
194    Ok(operations::serialize::vector_serialize(vector))
195}
196
197pub fn vector_slice(args: &[Register]) -> Result<Value> {
198    if args.len() != 3 {
199        return Err(LimboError::InvalidArgument(
200            "vector_slice requires exactly three arguments".into(),
201        ));
202    }
203    let value_0 = args[0].get_value();
204    let value_1 = args[1].get_value().as_value_ref();
205    let value_2 = args[2].get_value().as_value_ref();
206
207    let vector = parse_vector(value_0, None)?;
208
209    let start_index = value_1
210        .as_int()
211        .ok_or_else(|| LimboError::InvalidArgument("start index must be an integer".into()))?;
212
213    let end_index = value_2
214        .as_int()
215        .ok_or_else(|| LimboError::InvalidArgument("end_index must be an integer".into()))?;
216
217    if start_index < 0 || end_index < 0 {
218        return Err(LimboError::InvalidArgument(
219            "start index and end_index must be non-negative".into(),
220        ));
221    }
222
223    let result =
224        operations::slice::vector_slice(&vector, start_index as usize, end_index as usize)?;
225
226    Ok(operations::serialize::vector_serialize(result))
227}