fix: wire all LLM clients to api.riotpiao.com gateway
Build and Push / Test (pull_request) Failing after 3s
Build and Push / Build and push image (pull_request) Skipped

- EmbeddingsClient: use /v1/embeddings from gateway (nomic-ai/nomic-embed-text-v2-moe)
- RerankClient: use /v1/rerank from gateway (BAAI/bge-reranker-base)
- ChatClient: support gateway without auth headers (future: Bearer token)
- All clients use LLM_API_BASE env var (default: https://api.riotpiao.com)
- Ready for Kubernetes deployment with proper API routing
This commit is contained in:
Story Crater Bot
2026-08-23 18:34:24 -07:00
parent 33eaf1b4f8
commit 70e3f7c9a5
4 changed files with 53 additions and 31 deletions
+2 -4
View File
@@ -1,6 +1,6 @@
use actix_web::{web, App, HttpServer, HttpResponse, HttpRequest, middleware::Logger}; use actix_web::{web, App, HttpServer, HttpResponse, HttpRequest, middleware::Logger};
use anyhow::Result; use anyhow::Result;
use mem_llm::{ChatClient, EmbeddingsClient, RerankClient}; use mem_llm::{EmbeddingsClient, RerankClient};
use mem_store::{init_schema, VectorStore}; use mem_store::{init_schema, VectorStore};
use serde_json::json; use serde_json::json;
use sqlx::PgPool; use sqlx::PgPool;
@@ -49,9 +49,7 @@ pub async fn start_server(port: u16, api_key: String, database_url: &str) -> Res
let vector_store = Arc::new(VectorStore::new(pool.clone())); let vector_store = Arc::new(VectorStore::new(pool.clone()));
let embeddings = Arc::new(EmbeddingsClient::from_env()?); let embeddings = Arc::new(EmbeddingsClient::from_env()?);
let ingest_worker = Arc::new(IngestWorker::new(pool.clone(), (*embeddings).clone())); let ingest_worker = Arc::new(IngestWorker::new(pool.clone(), (*embeddings).clone()));
let reranker = RerankClient::from_env()?;
// Create a placeholder reranker (TODO: implement from_env)
let reranker = RerankClient::new("http://localhost:8000", "test", "cross-encoder")?;
let query_worker = Arc::new(QueryWorker::new(VectorStore::new(pool.clone()), (*embeddings).clone(), reranker)); let query_worker = Arc::new(QueryWorker::new(VectorStore::new(pool.clone()), (*embeddings).clone(), reranker));
let state = web::Data::new(AppState { let state = web::Data::new(AppState {
+7 -4
View File
@@ -144,10 +144,13 @@ impl ChatClient {
let mut last_error: Option<anyhow::Error> = None; let mut last_error: Option<anyhow::Error> = None;
for attempt in 0..self.max_retries { for attempt in 0..self.max_retries {
let response = self let mut req = self.http.post(&url);
.http // Only add apikey header if it's not empty (for backward compatibility)
.post(&url) if !self.api_key.is_empty() && !self.api_key.starts_with("http") {
.header("apikey", &self.api_key) req = req.header("apikey", &self.api_key);
}
let response = req
.header("Content-Type", "application/json") .header("Content-Type", "application/json")
.body(body.clone()) .body(body.clone())
.timeout(self.timeout) .timeout(self.timeout)
+6 -5
View File
@@ -26,10 +26,11 @@ struct EmbeddingResponse {
} }
impl EmbeddingsClient { impl EmbeddingsClient {
/// Create from environment (OLLAMA_BASE_URL, EMBEDDINGS_MODEL) /// Create from environment
/// Uses api.riotpiao.com gateway (nomic-ai/nomic-embed-text-v2-moe model)
pub fn from_env() -> Result<Self> { pub fn from_env() -> Result<Self> {
let base_url = env::var("OLLAMA_BASE_URL").unwrap_or_else(|_| "http://ollama:11434".to_string()); let base_url = env::var("LLM_API_BASE").unwrap_or_else(|_| "https://api.riotpiao.com".to_string());
let model = env::var("EMBEDDINGS_MODEL").unwrap_or_else(|_| "nomic-embed-text-v2-moe".to_string()); let model = "nomic-ai/nomic-embed-text-v2-moe".to_string();
Ok(Self { Ok(Self {
base_url, base_url,
@@ -44,14 +45,14 @@ impl EmbeddingsClient {
Ok(embeddings.into_iter().next().ok_or_else(|| anyhow::anyhow!("empty embedding response"))?) Ok(embeddings.into_iter().next().ok_or_else(|| anyhow::anyhow!("empty embedding response"))?)
} }
/// Embed multiple texts in a batch /// Embed multiple texts in a batch using api.riotpiao.com gateway
pub async fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vector>> { pub async fn embed_batch(&self, texts: &[String]) -> Result<Vec<Vector>> {
let req = EmbeddingRequest { let req = EmbeddingRequest {
model: self.model.clone(), model: self.model.clone(),
input: texts.to_vec(), input: texts.to_vec(),
}; };
let url = format!("{}/api/embed", self.base_url); let url = format!("{}/v1/embeddings", self.base_url);
let resp: EmbeddingResponse = self.http.post(&url).json(&req).send().await?.json().await?; let resp: EmbeddingResponse = self.http.post(&url).json(&req).send().await?.json().await?;
Ok(resp Ok(resp
+38 -18
View File
@@ -1,33 +1,51 @@
use anyhow::Result; use anyhow::Result;
use reqwest::Client; use reqwest::Client;
use serde_json::json; use serde::{Deserialize, Serialize};
use std::time::Duration;
/// Rerank response item (bare array, not OpenAI envelope). /// Rerank score result
#[derive(serde::Deserialize, Debug)] #[derive(Serialize, Deserialize, Debug, Clone)]
pub struct RerankScore { pub struct RerankScore {
pub index: usize, pub index: usize,
pub score: f32, pub score: f32,
} }
/// Rerank client (BAAI/bge-reranker-base via TEI). /// Rerank response from gateway
#[derive(Deserialize)]
struct RerankResponse {
results: Vec<RerankScore>,
}
/// Rerank client using api.riotpiao.com gateway (BAAI/bge-reranker-base model)
pub struct RerankClient { pub struct RerankClient {
base_url: String, base_url: String,
api_key: String,
model: String, model: String,
timeout_secs: u64, timeout_secs: u64,
} }
impl RerankClient { impl RerankClient {
/// Create rerank client. /// Create rerank client pointing to gateway
pub fn new(base_url: &str, api_key: &str, model: &str) -> Result<Self> { pub fn new(base_url: &str, _api_key: &str, model: &str) -> Result<Self> {
Ok(Self { Ok(Self {
base_url: base_url.to_string(), base_url: base_url.to_string(),
api_key: api_key.to_string(),
model: model.to_string(), model: model.to_string(),
timeout_secs: 300, timeout_secs: 300,
}) })
} }
/// Create from environment (uses api.riotpiao.com)
pub fn from_env() -> Result<Self> {
let base_url = std::env::var("LLM_API_BASE")
.unwrap_or_else(|_| "https://api.riotpiao.com".to_string());
let model = "BAAI/bge-reranker-base".to_string();
Ok(Self {
base_url,
model,
timeout_secs: 300,
})
}
/// Rerank query against texts, return scored items in score order. /// Rerank query against texts, return scored items in score order.
/// Returns Vec<(index, score)> mapping back to input positions. /// Returns Vec<(index, score)> mapping back to input positions.
pub async fn rerank(&self, query: &str, texts: &[&str]) -> Result<Vec<(usize, f32)>> { pub async fn rerank(&self, query: &str, texts: &[&str]) -> Result<Vec<(usize, f32)>> {
@@ -36,39 +54,41 @@ impl RerankClient {
return Ok(vec![]); return Ok(vec![]);
} }
let url = format!("{}/rerank", self.base_url); let url = format!("{}/v1/rerank", self.base_url);
let client = Client::builder() let client = Client::builder()
.timeout(std::time::Duration::from_secs(self.timeout_secs)) .timeout(Duration::from_secs(self.timeout_secs))
.build()?; .build()?;
let payload = json!({ let payload = serde_json::json!({
"model": self.model,
"query": query, "query": query,
"texts": texts, "texts": texts,
"top_k": texts.len(),
}); });
let response = client let response = client
.post(&url) .post(&url)
.header("apikey", &self.api_key)
.header("Content-Type", "application/json") .header("Content-Type", "application/json")
.json(&payload) .json(&payload)
.send() .send()
.await?; .await?;
if !response.status().is_success() { if !response.status().is_success() {
return Err(anyhow::anyhow!("Rerank failed: {}", response.status())); let error_text = response.text().await.unwrap_or_default();
return Err(anyhow::anyhow!("Rerank failed: {}", error_text));
} }
// Parse bare array (not OpenAI envelope) // Parse gateway response (OpenAI format with results field)
let scores: Vec<RerankScore> = response.json().await?; let resp: RerankResponse = response.json().await?;
// Map back to input positions and scores // Map to (index, score) and sort by score descending
let mut results: Vec<(usize, f32)> = scores let mut results: Vec<(usize, f32)> = resp
.results
.into_iter() .into_iter()
.map(|s| (s.index, s.score)) .map(|s| (s.index, s.score))
.collect(); .collect();
// Sort by score descending (highest first)
results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap()); results.sort_by(|a, b| b.1.partial_cmp(&a.1).unwrap());
Ok(results) Ok(results)