tidaldb/applications/forage/embedder/src/main.rs
jordan c1c5a10fbc chore: reuse reqwest::Client across requests in forage embedder; minor forage updates
Co-Authored-By: Claude Sonnet 4.6 <noreply@anthropic.com>
2026-02-23 23:16:32 -07:00

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");
}