Browse Source

feat(jina): add multi-key failover support for embedding, reranking and wiki sync

master
nwhco 3 weeks ago
parent
commit
9ca93ca366
  1. 3
      config/production.env
  2. 38
      src/knowledge/embedding_factory.py
  3. 163
      src/knowledge/multi_jina_embedder.py
  4. 37
      src/knowledge/sync_wiki.py
  5. 21
      src/utils/jina_keys.py
  6. 157
      src/utils/reranker.py

3
config/production.env

@ -35,7 +35,8 @@ BASE_COLLECTION_NAME=imam_reza_collection
EMBEDDER_MODEL=all-MiniLM-L6-v2
RERANKER_MODEL=jina-reranker-v3
EMBEDDER_DIMENSIONS=384
JINA_API_KEY=jina_acfe129adad644c494e085b386736d72kXVyA_ZdcbHq6nRb1GrHPkqb8bw9
JINA_API_KEYS=jina_a72e688ae2b6476abef817e5426bb28dh5KTQKAk_4ydErx9AfZSvZVpcPj_,jina_69cf5ef72f2c44ad85d27f85283585d4wtH3VQQVNsV9goLpnps5rpE4L_Gf,jina_64b3c101b0be43ef9372f4282c73cec3TXZsQz_Slwbf0BVA0xZ2mQAWF-cC
JINA_API_KEY=jina_a72e688ae2b6476abef817e5426bb28dh5KTQKAk_4ydErx9AfZSvZVpcPj_
OPENAI_API_KEY=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
# ---------------- LANGFUSE Settings ----------------

38
src/knowledge/embedding_factory.py

@ -1,67 +1,55 @@
import yaml
import os
from typing import Optional
from agno.knowledge.embedder.openai import OpenAIEmbedder
from agno.knowledge.embedder.jina import JinaEmbedder
from pathlib import Path
# If Agno supports generic OpenAI-like embedders, we use OpenAIEmbedder with base_url
from agno.knowledge.embedder.openai import OpenAIEmbedder
from src.knowledge.multi_jina_embedder import MultiKeyJinaEmbedder
from src.utils.jina_keys import get_jina_keys
class EmbeddingFactory:
def __init__(self):
# Get the directory where this file (factory.py) is located
current_file_path = Path(__file__).resolve()
# Navigate up to the project root
# If structure is: /app/src/models/factory.py
# .parent = models, .parent = src, .parent = app (root)
project_root = current_file_path.parent.parent.parent
# Construct the absolute path
config_path = project_root / 'config' / 'embeddings.yaml'
print(f"Loading config from: {config_path}") # Debug log
with open(config_path) as f:
# Simple variable expansion for ${VAR}
print(f"Loading config from: {config_path}")
with open(config_path, "r", encoding="utf-8") as f:
content = f.read()
for key, val in os.environ.items():
content = content.replace(f"${{{key}}}", val)
self.config = yaml.safe_load(content)
def get_embedder(self, model_name: Optional[str] = None):
# 1. Default Logic
if model_name is None:
model_name = self.config['embeddings']['default']
models_config = self.config['embeddings']['models']
if model_name not in models_config:
raise ValueError(f"Embedding model '{model_name}' not found in config.")
config = models_config[model_name]
provider = config['provider']
# # 2. Provider Logic
api_key_env = config.get('api_key')
if api_key_env and api_key_env.startswith("${"):
if api_key_env and str(api_key_env).startswith("${"):
api_key = os.getenv(api_key_env[2:-1])
else:
api_key = api_key_env
# CASE B: OpenAI (Official)
if provider == "openai":
return OpenAIEmbedder(
id=config['id'],
dimensions=config['dimensions'],
api_key=api_key
)
# CASE C: OpenAI Compatible (Jina API, etc.)
elif provider == "jinaai":
return JinaEmbedder(
all_keys = get_jina_keys()
if api_key and api_key not in all_keys:
all_keys.insert(0, api_key)
return MultiKeyJinaEmbedder(
id=config['id'],
dimensions=config['dimensions'],
api_key=api_key
api_keys=all_keys,
api_key=all_keys[0] if all_keys else None
)
print(f"Unknown provider type: {provider}")
raise ValueError(f"Unknown provider type: {provider}")
raise ValueError(f"Unknown provider type: {provider}")

163
src/knowledge/multi_jina_embedder.py

@ -0,0 +1,163 @@
import os
import requests
import aiohttp
import logging
from dataclasses import dataclass, field
from typing import List, Optional, Dict, Any, Tuple
from agno.knowledge.embedder.jina import JinaEmbedder
from src.utils.jina_keys import get_jina_keys
logger = logging.getLogger(__name__)
@dataclass
class MultiKeyJinaEmbedder(JinaEmbedder):
"""
JinaEmbedder subclass that supports multiple API keys with automatic failover/rotation.
If a key hits balance limit (402), forbidden (403), unauthorized (401), or rate limit (429),
it automatically switches to the next available key and retries.
"""
api_keys: List[str] = field(default_factory=get_jina_keys)
_current_key_idx: int = field(default=0, init=False, repr=False)
def __post_init__(self):
if not self.api_keys:
self.api_keys = get_jina_keys()
if self.api_key and self.api_key not in self.api_keys:
self.api_keys.insert(0, self.api_key)
if self.api_keys:
self.api_key = self.api_keys[0]
def _get_active_key(self) -> str:
if not self.api_keys:
if self.api_key:
return self.api_key
raise ValueError("No Jina API keys provided in JINA_API_KEYS or JINA_API_KEY")
return self.api_keys[self._current_key_idx % len(self.api_keys)]
def _rotate_key(self) -> str:
if len(self.api_keys) > 1:
prev = self._get_active_key()
self._current_key_idx = (self._current_key_idx + 1) % len(self.api_keys)
new_key = self._get_active_key()
self.api_key = new_key
logger.warning(f"🔄 Rotating Jina API Key from {prev[:12]}... to {new_key[:12]}...")
return new_key
return self._get_active_key()
def _get_headers_for_key(self, key: str) -> Dict[str, str]:
headers = {"Content-Type": "application/json", "Authorization": f"Bearer {key}"}
if self.headers:
headers.update(self.headers)
return headers
def _response(self, text: str) -> Dict[str, Any]:
data = {
"model": self.id,
"late_chunking": self.late_chunking,
"dimensions": self.dimensions,
"embedding_type": self.embedding_type,
"input": [text],
}
if self.user is not None:
data["user"] = self.user
if self.request_params:
data.update(self.request_params)
keys_to_try = max(len(self.api_keys), 1)
last_exception = None
for _ in range(keys_to_try):
key = self._get_active_key()
try:
headers = self._get_headers_for_key(key)
response = requests.post(self.base_url, headers=headers, json=data, timeout=self.timeout or 30.0)
if response.status_code in (401, 402, 403, 429):
logger.warning(f"⚠️ Jina key {key[:12]}... failed with HTTP {response.status_code}: {response.text[:120]}")
self._rotate_key()
continue
response.raise_for_status()
return response.json()
except Exception as e:
last_exception = e
logger.warning(f"⚠️ Exception with Jina key {key[:12]}...: {e}. Trying next key...")
self._rotate_key()
if last_exception:
raise last_exception
raise RuntimeError("All Jina API keys failed")
async def _async_response(self, text: str) -> Dict[str, Any]:
data = {
"model": self.id,
"late_chunking": self.late_chunking,
"dimensions": self.dimensions,
"embedding_type": self.embedding_type,
"input": [text],
}
if self.user is not None:
data["user"] = self.user
if self.request_params:
data.update(self.request_params)
timeout = aiohttp.ClientTimeout(total=self.timeout or 30.0)
keys_to_try = max(len(self.api_keys), 1)
last_exception = None
for _ in range(keys_to_try):
key = self._get_active_key()
try:
headers = self._get_headers_for_key(key)
async with aiohttp.ClientSession(timeout=timeout) as session:
async with session.post(self.base_url, headers=headers, json=data) as response:
if response.status in (401, 402, 403, 429):
logger.warning(f"⚠️ Async Jina key {key[:12]}... failed with HTTP {response.status}")
self._rotate_key()
continue
response.raise_for_status()
return await response.json()
except Exception as e:
last_exception = e
logger.warning(f"⚠️ Async Jina exception with key {key[:12]}...: {e}. Trying next key...")
self._rotate_key()
if last_exception:
raise last_exception
raise RuntimeError("All Jina API keys failed in async request")
async def _async_batch_response(self, texts: List[str]) -> Dict[str, Any]:
data = {
"model": self.id,
"late_chunking": self.late_chunking,
"dimensions": self.dimensions,
"embedding_type": self.embedding_type,
"input": texts,
}
if self.user is not None:
data["user"] = self.user
if self.request_params:
data.update(self.request_params)
timeout = aiohttp.ClientTimeout(total=self.timeout or 60.0)
keys_to_try = max(len(self.api_keys), 1)
last_exception = None
for _ in range(keys_to_try):
key = self._get_active_key()
try:
headers = self._get_headers_for_key(key)
async with aiohttp.ClientSession(timeout=timeout) as session:
async with session.post(self.base_url, headers=headers, json=data) as response:
if response.status in (401, 402, 403, 429):
logger.warning(f"⚠️ Async batch Jina key {key[:12]}... failed with HTTP {response.status}")
self._rotate_key()
continue
response.raise_for_status()
return await response.json()
except Exception as e:
last_exception = e
logger.warning(f"⚠️ Async batch Jina exception with key {key[:12]}...: {e}. Trying next key...")
self._rotate_key()
if last_exception:
raise last_exception
raise RuntimeError("All Jina API keys failed in async batch request")

37
src/knowledge/sync_wiki.py

@ -29,22 +29,27 @@ def get_text_from_json(json_data, target_lang='fa'):
return "Unknown"
def convert_html_to_md_jina(html_content: str, row_id: int) -> str:
"""Helper to call Jina AI and convert HTML to clean Markdown."""
headers = {
"Authorization": f"Bearer {JINA_API_KEY}",
"Accept": "application/json"
}
files = {
'file': (f'document_{row_id}.html', html_content, 'text/html')
}
try:
response = requests.post("https://r.jina.ai/", headers=headers, files=files, timeout=30)
response.raise_for_status()
jina_data = response.json().get('data', {})
return jina_data.get('content', '')
except Exception as e:
print(f"⚠️ [Jina Error] ID {row_id}: {e}")
return ""
"""Helper to call Jina AI and convert HTML to clean Markdown with multi-key failover."""
from src.utils.jina_keys import get_jina_keys
keys = get_jina_keys()
for key in keys:
headers = {
"Authorization": f"Bearer {key}",
"Accept": "application/json"
}
files = {
'file': (f'document_{row_id}.html', html_content, 'text/html')
}
try:
response = requests.post("https://r.jina.ai/", headers=headers, files=files, timeout=30)
if response.status_code in (401, 402, 403, 429):
continue
response.raise_for_status()
jina_data = response.json().get('data', {})
return jina_data.get('content', '')
except Exception as e:
print(f"⚠️ [Jina Error with key {key[:12]}] ID {row_id}: {e}")
return ""
def run_wiki_embedding_sync(session_id: int):

21
src/utils/jina_keys.py

@ -0,0 +1,21 @@
import os
from typing import List
def get_jina_keys() -> List[str]:
"""
Returns a list of Jina API keys from environment variables.
Supports both JINA_API_KEYS (comma-separated) and JINA_API_KEY.
"""
keys = []
keys_env = os.getenv("JINA_API_KEYS", "")
if keys_env:
for k in keys_env.split(","):
cleaned = k.strip()
if cleaned and cleaned not in keys:
keys.append(cleaned)
single_key = os.getenv("JINA_API_KEY", "").strip()
if single_key and single_key not in keys:
keys.append(single_key)
return keys

157
src/utils/reranker.py

@ -1,96 +1,19 @@
# import os
# import requests
# import json
# from typing import List, Any
# from dotenv import load_dotenv
# load_dotenv()
# def rerank_documents(query: str, documents: List[Any], top_n: int = 3) -> List[Any]:
# """
# Reranks a list of documents using Jina AI's Reranker API.
# Args:
# query: The user's question.
# documents: List of document objects (must have a .content attribute).
# top_n: How many top documents to return.
# Returns:
# The top_n sorted document objects.
# """
# print(f"🔍🔍🔍🔍🔍 Reranking documents🔍🔍🔍🔍🔍")
# api_key = os.getenv("JINA_API_KEY")
# if not api_key:
# print("⚠️ JINA_API_KEY not found. Returning original order.")
# return documents[:top_n]
# # 1. Prepare data for Jina
# # Jina needs a list of strings. We extract .content from your Agno Document objects.
# doc_contents = [doc.content for doc in documents]
# url = "https://api.jina.ai/v1/rerank"
# headers = {
# "Content-Type": "application/json",
# "Authorization": f"Bearer {api_key}"
# }
# payload = {
# "model": "jina-reranker-v3", # Best for mixed language (English/Arabic/Persian)
# "query": query,
# "documents": doc_contents,
# "top_n": top_n
# }
# try:
# # 2. Call Jina API
# response = requests.post(url, headers=headers, json=payload)
# response.raise_for_status()
# results = response.json()["results"]
# # 3. Map back to original Document objects
# # Jina returns indices (e.g., "index 4 is the best"). We use these to pick from your original list.
# reranked_docs = []
# for result in results:
# original_index = result["index"]
# relevance_score = result["relevance_score"]
# doc = documents[original_index]
# # 👇 FIX 1: Ensure meta_data exists before writing to it
# if not hasattr(doc, "meta_data") or doc.meta_data is None:
# doc.meta_data = {}
# # 👇 FIX 2: Use .meta_data (with underscore)
# doc.meta_data["rerank_score"] = relevance_score
# reranked_docs.append(doc)
# print(f"✨ Reranked {len(documents)} docs -> Top {len(reranked_docs)}")
# return reranked_docs
# except Exception as e:
# print(f"❌ Reranking failed: {e}. Falling back to vector search order.")
# return documents[:top_n]
import os
import yaml
import requests
import re
from typing import List, Any, Dict, Optional
from dotenv import load_dotenv
from src.utils.jina_keys import get_jina_keys
load_dotenv()
class Reranker:
def __init__(self, config_path: str = "config/rerankers.yaml"):
self.config = self._load_config(config_path)
# 1. Get the active model configuration
self.active_model_name = self.config["rerankers"]["default"]
self.model_config = self.config["rerankers"]["models"][self.active_model_name]
# 2. Extract key params
self.provider = self.model_config.get("provider")
self.model_id = self.model_config.get("model")
self.default_top_n = self.model_config.get("top_n", 3)
@ -100,14 +23,12 @@ class Reranker:
print(f"🚀 Initialized Reranker: {self.active_model_name} ({self.provider})")
def _load_config(self, path: str) -> Dict:
"""Loads YAML and replaces ${VAR} with env variables."""
if not os.path.exists(path):
raise FileNotFoundError(f"Config file not found at: {path}")
with open(path, "r", encoding="utf-8") as f:
content = f.read()
# Regex to find ${VAR_NAME} and replace with os.getenv('VAR_NAME')
pattern = re.compile(r'\$\{(\w+)\}')
def replace(match):
env_var = match.group(1)
@ -117,9 +38,6 @@ class Reranker:
return yaml.safe_load(updated_content)
def rerank_documents(self, query: str, documents: List[Any], top_n: Optional[int] = None) -> List[Any]:
"""
Main entry point for reranking.
"""
final_top_n = top_n if top_n is not None else self.default_top_n
if not documents:
@ -128,7 +46,6 @@ class Reranker:
print(f"🔍 Reranking {len(documents)} docs using {self.provider}...")
try:
# Route to the correct provider logic
if self.provider == "jinaai":
return self._rerank_jina(query, documents, final_top_n)
else:
@ -139,17 +56,15 @@ class Reranker:
return documents[:final_top_n]
def _rerank_jina(self, query: str, documents: List[Any], top_n: int) -> List[Any]:
if not self.api_key:
keys = get_jina_keys()
if self.api_key and self.api_key not in keys:
keys.insert(0, self.api_key)
if not keys:
print("⚠️ Missing Jina API Key. Skipping rerank.")
return documents[:top_n]
# Prepare payload
doc_contents = [getattr(doc, "content", str(doc)) for doc in documents]
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {self.api_key}"
}
payload = {
"model": self.model_id,
"query": query,
@ -157,33 +72,39 @@ class Reranker:
"top_n": top_n
}
response = requests.post(self.base_url, headers=headers, json=payload)
response.raise_for_status()
results = response.json()["results"]
# Map results back to original documents
reranked_docs = []
for result in results:
original_index = result["index"]
relevance_score = result["relevance_score"]
doc = documents[original_index]
# Safely add metadata
if not hasattr(doc, "meta_data") or doc.meta_data is None:
doc.meta_data = {}
doc.meta_data["rerank_score"] = relevance_score
doc.meta_data["rerank_model"] = self.model_id
reranked_docs.append(doc)
print(f"✨ Top score: {results[0]['relevance_score']}")
return reranked_docs
# Singleton instance to avoid reloading config on every request
for key in keys:
try:
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {key}"
}
response = requests.post(self.base_url, headers=headers, json=payload, timeout=30)
if response.status_code in (401, 402, 403, 429):
print(f"⚠️ Jina Rerank key {key[:12]}... failed with {response.status_code}. Trying next key...")
continue
response.raise_for_status()
results = response.json().get("results", [])
reranked_docs = []
for result in results:
original_index = result["index"]
relevance_score = result["relevance_score"]
doc = documents[original_index]
if not hasattr(doc, "meta_data") or doc.meta_data is None:
doc.meta_data = {}
doc.meta_data["rerank_score"] = relevance_score
doc.meta_data["rerank_model"] = self.model_id
reranked_docs.append(doc)
print(f"✨ Top score: {results[0]['relevance_score']} (key: {key[:12]}...)")
return reranked_docs
except Exception as e:
print(f"⚠️ Jina Rerank exception with key {key[:12]}...: {e}. Trying next...")
print("❌ All Jina Rerank keys failed. Falling back to vector search order.")
return documents[:top_n]
reranker_instance = Reranker()
# Public function interface (keeps your existing code working)
def rerank_documents(query: str, documents: List[Any], top_n: int = 3) -> List[Any]:
return reranker_instance.rerank_documents(query, documents, top_n)
return reranker_instance.rerank_documents(query, documents, top_n)
Loading…
Cancel
Save