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. 36
      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. 153
      src/utils/reranker.py

3
config/production.env

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

36
src/knowledge/embedding_factory.py

@ -1,67 +1,55 @@
import yaml import yaml
import os import os
from typing import Optional from typing import Optional
from agno.knowledge.embedder.openai import OpenAIEmbedder
from agno.knowledge.embedder.jina import JinaEmbedder
from pathlib import Path 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: class EmbeddingFactory:
def __init__(self): def __init__(self):
# Get the directory where this file (factory.py) is located
current_file_path = Path(__file__).resolve() 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 project_root = current_file_path.parent.parent.parent
# Construct the absolute path
config_path = project_root / 'config' / 'embeddings.yaml' 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() content = f.read()
for key, val in os.environ.items(): for key, val in os.environ.items():
content = content.replace(f"${{{key}}}", val) content = content.replace(f"${{{key}}}", val)
self.config = yaml.safe_load(content) self.config = yaml.safe_load(content)
def get_embedder(self, model_name: Optional[str] = None): def get_embedder(self, model_name: Optional[str] = None):
# 1. Default Logic
if model_name is None: if model_name is None:
model_name = self.config['embeddings']['default'] model_name = self.config['embeddings']['default']
models_config = self.config['embeddings']['models'] models_config = self.config['embeddings']['models']
if model_name not in models_config: if model_name not in models_config:
raise ValueError(f"Embedding model '{model_name}' not found in config.") raise ValueError(f"Embedding model '{model_name}' not found in config.")
config = models_config[model_name] config = models_config[model_name]
provider = config['provider'] provider = config['provider']
# # 2. Provider Logic
api_key_env = config.get('api_key') 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]) api_key = os.getenv(api_key_env[2:-1])
else: else:
api_key = api_key_env api_key = api_key_env
# CASE B: OpenAI (Official)
if provider == "openai": if provider == "openai":
return OpenAIEmbedder( return OpenAIEmbedder(
id=config['id'], id=config['id'],
dimensions=config['dimensions'], dimensions=config['dimensions'],
api_key=api_key api_key=api_key
) )
# CASE C: OpenAI Compatible (Jina API, etc.)
elif provider == "jinaai": 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'], id=config['id'],
dimensions=config['dimensions'], 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" return "Unknown"
def convert_html_to_md_jina(html_content: str, row_id: int) -> str: 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): 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

153
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 os
import yaml import yaml
import requests import requests
import re import re
from typing import List, Any, Dict, Optional from typing import List, Any, Dict, Optional
from dotenv import load_dotenv from dotenv import load_dotenv
from src.utils.jina_keys import get_jina_keys
load_dotenv() load_dotenv()
class Reranker: class Reranker:
def __init__(self, config_path: str = "config/rerankers.yaml"): def __init__(self, config_path: str = "config/rerankers.yaml"):
self.config = self._load_config(config_path) self.config = self._load_config(config_path)
# 1. Get the active model configuration
self.active_model_name = self.config["rerankers"]["default"] self.active_model_name = self.config["rerankers"]["default"]
self.model_config = self.config["rerankers"]["models"][self.active_model_name] self.model_config = self.config["rerankers"]["models"][self.active_model_name]
# 2. Extract key params
self.provider = self.model_config.get("provider") self.provider = self.model_config.get("provider")
self.model_id = self.model_config.get("model") self.model_id = self.model_config.get("model")
self.default_top_n = self.model_config.get("top_n", 3) 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})") print(f"🚀 Initialized Reranker: {self.active_model_name} ({self.provider})")
def _load_config(self, path: str) -> Dict: def _load_config(self, path: str) -> Dict:
"""Loads YAML and replaces ${VAR} with env variables."""
if not os.path.exists(path): if not os.path.exists(path):
raise FileNotFoundError(f"Config file not found at: {path}") raise FileNotFoundError(f"Config file not found at: {path}")
with open(path, "r", encoding="utf-8") as f: with open(path, "r", encoding="utf-8") as f:
content = f.read() content = f.read()
# Regex to find ${VAR_NAME} and replace with os.getenv('VAR_NAME')
pattern = re.compile(r'\$\{(\w+)\}') pattern = re.compile(r'\$\{(\w+)\}')
def replace(match): def replace(match):
env_var = match.group(1) env_var = match.group(1)
@ -117,9 +38,6 @@ class Reranker:
return yaml.safe_load(updated_content) return yaml.safe_load(updated_content)
def rerank_documents(self, query: str, documents: List[Any], top_n: Optional[int] = None) -> List[Any]: 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 final_top_n = top_n if top_n is not None else self.default_top_n
if not documents: if not documents:
@ -128,7 +46,6 @@ class Reranker:
print(f"🔍 Reranking {len(documents)} docs using {self.provider}...") print(f"🔍 Reranking {len(documents)} docs using {self.provider}...")
try: try:
# Route to the correct provider logic
if self.provider == "jinaai": if self.provider == "jinaai":
return self._rerank_jina(query, documents, final_top_n) return self._rerank_jina(query, documents, final_top_n)
else: else:
@ -139,17 +56,15 @@ class Reranker:
return documents[:final_top_n] return documents[:final_top_n]
def _rerank_jina(self, query: str, documents: List[Any], top_n: int) -> List[Any]: 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.") print("⚠️ Missing Jina API Key. Skipping rerank.")
return documents[:top_n] return documents[:top_n]
# Prepare payload
doc_contents = [getattr(doc, "content", str(doc)) for doc in documents] doc_contents = [getattr(doc, "content", str(doc)) for doc in documents]
headers = {
"Content-Type": "application/json",
"Authorization": f"Bearer {self.api_key}"
}
payload = { payload = {
"model": self.model_id, "model": self.model_id,
"query": query, "query": query,
@ -157,33 +72,39 @@ class Reranker:
"top_n": top_n "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
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]
# Singleton instance to avoid reloading config on every request
reranker_instance = Reranker() 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]: 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