210 lines
5.9 KiB
Rust
210 lines
5.9 KiB
Rust
/// Forage Embedding Sidecar
|
|
///
|
|
/// A lightweight HTTP server that accepts text and returns a float vector.
|
|
/// The vector is suitable for ANN retrieval in tidalDB's HNSW index.
|
|
///
|
|
/// Modes:
|
|
/// --mock Deterministic pseudo-random unit vector derived from FNV-1a hash
|
|
/// of the input text. No API key required. Useful for development
|
|
/// and architecture testing; vectors are stable across runs but carry
|
|
/// no semantic meaning.
|
|
///
|
|
/// (default) OpenAI text-embedding-3-small via OPENAI_API_KEY env var.
|
|
/// Requires a valid key. Produces genuine semantic vectors.
|
|
///
|
|
/// Usage:
|
|
/// forage-embedder --mock # dev / no key
|
|
/// OPENAI_API_KEY=sk-... forage-embedder # production
|
|
///
|
|
/// Protocol:
|
|
/// POST /embed
|
|
/// Body: { "text": "..." }
|
|
/// Reply: { "vector": [f32, ...], "dim": 1536 }
|
|
use std::net::SocketAddr;
|
|
use std::sync::Arc;
|
|
|
|
use axum::Json;
|
|
use axum::Router;
|
|
use axum::extract::State;
|
|
use axum::http::StatusCode;
|
|
use axum::response::IntoResponse;
|
|
use axum::routing::post;
|
|
use clap::Parser;
|
|
use serde::{Deserialize, Serialize};
|
|
use tower_http::cors::CorsLayer;
|
|
|
|
const DIM: usize = 1536;
|
|
|
|
#[derive(Parser)]
|
|
#[command(
|
|
name = "forage-embedder",
|
|
about = "Forage embedding sidecar (mock or OpenAI)"
|
|
)]
|
|
struct Args {
|
|
/// Use deterministic mock embeddings (no API key required).
|
|
#[arg(long)]
|
|
mock: bool,
|
|
|
|
/// Port to listen on.
|
|
#[arg(long, default_value = "4243")]
|
|
port: u16,
|
|
}
|
|
|
|
#[derive(Clone)]
|
|
enum Mode {
|
|
Mock,
|
|
/// `client` is created once at startup and reused across requests.
|
|
/// `reqwest::Client` is cheaply cloneable (`Arc`-backed connection pool).
|
|
OpenAi {
|
|
api_key: String,
|
|
client: reqwest::Client,
|
|
},
|
|
}
|
|
|
|
#[derive(Deserialize)]
|
|
struct EmbedReq {
|
|
text: String,
|
|
}
|
|
|
|
#[derive(Serialize)]
|
|
struct EmbedResp {
|
|
vector: Vec<f32>,
|
|
dim: usize,
|
|
}
|
|
|
|
async fn post_embed(State(mode): State<Arc<Mode>>, Json(req): Json<EmbedReq>) -> impl IntoResponse {
|
|
let vector = match mode.as_ref() {
|
|
Mode::Mock => mock_embed(&req.text),
|
|
Mode::OpenAi { api_key, client } => match openai_embed(client, api_key, &req.text).await {
|
|
Ok(v) => v,
|
|
Err(e) => {
|
|
return (
|
|
StatusCode::BAD_GATEWAY,
|
|
Json(serde_json::json!({ "error": e.to_string() })),
|
|
)
|
|
.into_response();
|
|
}
|
|
},
|
|
};
|
|
(
|
|
StatusCode::OK,
|
|
Json(EmbedResp {
|
|
dim: vector.len(),
|
|
vector,
|
|
}),
|
|
)
|
|
.into_response()
|
|
}
|
|
|
|
/// Deterministic mock embedding: FNV-1a hash of text → seeded LCG → 1536-dim unit vector.
|
|
///
|
|
/// Properties:
|
|
/// - Same text → same vector (stable across runs)
|
|
/// - Different texts → different vectors (hash dispersion)
|
|
/// - No semantic meaning (random unit vectors)
|
|
fn mock_embed(text: &str) -> Vec<f32> {
|
|
// FNV-1a 64-bit hash as seed.
|
|
let mut state: u64 = 14_695_981_039_346_656_037;
|
|
for byte in text.bytes() {
|
|
state ^= u64::from(byte);
|
|
state = state.wrapping_mul(1_099_511_628_211);
|
|
}
|
|
|
|
let mut v = Vec::with_capacity(DIM);
|
|
for _ in 0..DIM {
|
|
// PCG-inspired step: fast, well-distributed.
|
|
state = state
|
|
.wrapping_mul(6_364_136_223_846_793_005)
|
|
.wrapping_add(1_442_695_040_888_963_407);
|
|
let sample = ((state >> 33) as f32 / u32::MAX as f32) * 2.0 - 1.0;
|
|
v.push(sample);
|
|
}
|
|
l2_normalize(&mut v);
|
|
v
|
|
}
|
|
|
|
/// OpenAI text-embedding-3-small call.
|
|
async fn openai_embed(
|
|
client: &reqwest::Client,
|
|
api_key: &str,
|
|
text: &str,
|
|
) -> Result<Vec<f32>, String> {
|
|
let resp = client
|
|
.post("https://api.openai.com/v1/embeddings")
|
|
.bearer_auth(api_key)
|
|
.json(&serde_json::json!({
|
|
"model": "text-embedding-3-small",
|
|
"input": text
|
|
}))
|
|
.send()
|
|
.await
|
|
.map_err(|e| format!("request failed: {e}"))?;
|
|
|
|
if !resp.status().is_success() {
|
|
let status = resp.status();
|
|
let body = resp.text().await.unwrap_or_default();
|
|
return Err(format!("OpenAI error {status}: {body}"));
|
|
}
|
|
|
|
let json: serde_json::Value = resp.json().await.map_err(|e| format!("parse error: {e}"))?;
|
|
|
|
let embedding = json["data"][0]["embedding"]
|
|
.as_array()
|
|
.ok_or("missing embedding field")?
|
|
.iter()
|
|
.map(|v| v.as_f64().unwrap_or(0.0) as f32)
|
|
.collect::<Vec<_>>();
|
|
|
|
Ok(embedding)
|
|
}
|
|
|
|
fn l2_normalize(v: &mut [f32]) {
|
|
let norm: f32 = v.iter().map(|x| x * x).sum::<f32>().sqrt();
|
|
if norm > 1e-9 {
|
|
for x in v.iter_mut() {
|
|
*x /= norm;
|
|
}
|
|
}
|
|
}
|
|
|
|
#[tokio::main]
|
|
async fn main() {
|
|
let args = Args::parse();
|
|
|
|
let mode = if args.mock {
|
|
println!("forage-embedder: mock mode (deterministic, no API key)");
|
|
Mode::Mock
|
|
} else {
|
|
let key = std::env::var("OPENAI_API_KEY").unwrap_or_else(|_| {
|
|
eprintln!(
|
|
"forage-embedder: OPENAI_API_KEY not set; falling back to mock mode.\n\
|
|
Set OPENAI_API_KEY or pass --mock to suppress this warning."
|
|
);
|
|
String::new()
|
|
});
|
|
if key.is_empty() {
|
|
Mode::Mock
|
|
} else {
|
|
println!("forage-embedder: OpenAI mode (text-embedding-3-small)");
|
|
Mode::OpenAi {
|
|
api_key: key,
|
|
client: reqwest::Client::new(),
|
|
}
|
|
}
|
|
};
|
|
|
|
let state = Arc::new(mode);
|
|
let app = Router::new()
|
|
.route("/embed", post(post_embed))
|
|
.layer(CorsLayer::permissive())
|
|
.with_state(state);
|
|
|
|
let addr = SocketAddr::from(([127, 0, 0, 1], args.port));
|
|
println!("forage-embedder listening on http://{addr}");
|
|
|
|
let listener = tokio::net::TcpListener::bind(addr)
|
|
.await
|
|
.expect("failed to bind");
|
|
axum::serve(listener, app).await.expect("server error");
|
|
}
|