S SmartDocs
Série: Ricky python 166 linhas · Atualizado 2026-04-07

providers.py

Ricky/RAG_project/llm/providers.py

"""
LLM and embedding provider management.
Handles health checks, creation, and automatic OpenRouter → Ollama failover.
"""
import json
import logging
from urllib.request import Request, urlopen
from urllib.error import URLError, HTTPError

from langchain_ollama import OllamaEmbeddings, OllamaLLM
from langchain_openai import ChatOpenAI

from config.settings import (
    OPENROUTER_API_KEY,
    OPENROUTER_BASE_URL,
    OPENROUTER_MODEL,
    OLLAMA_BASE_URL,
    OLLAMA_LLM_MODEL,
    EMBEDDING_MODEL,
)
from config.runtime import (
    get_llm_timeout,
    get_max_tokens,
    prefer_openrouter,
)

logger = logging.getLogger(__name__)


# ---------------------------------------------------------------------------
# Embeddings
# ---------------------------------------------------------------------------

def get_embeddings() -> OllamaEmbeddings:
    return OllamaEmbeddings(
        model=EMBEDDING_MODEL,
        base_url=OLLAMA_BASE_URL,
    )


def check_ollama_embedding_ready() -> bool:
    """Check that Ollama is running and the embedding model is available."""
    tags_url = f"{OLLAMA_BASE_URL}/api/tags"
    req = Request(tags_url, method="GET")
    try:
        with urlopen(req, timeout=8) as resp:
            payload = json.loads(resp.read().decode("utf-8"))
    except (URLError, HTTPError, TimeoutError) as e:
        logger.error("Cannot reach Ollama: %s", e)
        print(f"[ERROR] Cannot reach Ollama: {e}")
        print(f"[TIP] Start Ollama and confirm {OLLAMA_BASE_URL} is reachable.")
        return False

    models = [m.get("name", "") for m in payload.get("models", [])]
    if not any(
        name == EMBEDDING_MODEL or name.startswith(f"{EMBEDDING_MODEL}:")
        for name in models
    ):
        logger.error("Missing Ollama embedding model: %s", EMBEDDING_MODEL)
        print(f"[ERROR] Missing Ollama embedding model: {EMBEDDING_MODEL}")
        print(f"[TIP] Run: ollama pull {EMBEDDING_MODEL}")
        return False
    print("[OK] Ollama embeddings ready")
    return True


# ---------------------------------------------------------------------------
# LLM creation
# ---------------------------------------------------------------------------

def create_openrouter_llm(max_tokens: int | None = None, timeout: int | None = None) -> ChatOpenAI:
    return ChatOpenAI(
        model=OPENROUTER_MODEL,
        max_tokens=max_tokens or get_max_tokens(),
        temperature=0.1,
        timeout=timeout or get_llm_timeout(),
        api_key=OPENROUTER_API_KEY,
        base_url=OPENROUTER_BASE_URL,
    )


def create_ollama_llm(num_predict: int | None = None, timeout: int | None = None) -> OllamaLLM:
    return OllamaLLM(
        model=OLLAMA_LLM_MODEL,
        num_predict=num_predict or get_max_tokens(),
        temperature=0.1,
        sync_client_kwargs={"timeout": timeout or get_llm_timeout()},
    )


# ---------------------------------------------------------------------------
# Provider selection with health checks
# ---------------------------------------------------------------------------

def _model_name_matches(installed_name: str, target_name: str) -> bool:
    return installed_name == target_name


def check_openrouter_ready() -> bool:
    """Check that OpenRouter is reachable and the configured model is listed."""
    if not OPENROUTER_API_KEY:
        print("[ERROR] OPENROUTER_API_KEY is not set")
        print("[TIP] Add OPENROUTER_API_KEY=your_key to .env")
        return False

    models_url = f"{OPENROUTER_BASE_URL}/models"
    req = Request(
        models_url,
        method="GET",
        headers={"Authorization": f"Bearer {OPENROUTER_API_KEY}"},
    )
    try:
        with urlopen(req, timeout=8) as resp:
            payload = json.loads(resp.read().decode("utf-8"))
    except (URLError, HTTPError, TimeoutError) as e:
        logger.warning("Cannot reach OpenRouter: %s", e)
        print(f"[ERROR] Cannot reach OpenRouter: {e}")
        print("[TIP] Check OPENROUTER_BASE_URL and API key.")
        return False

    models = [m.get("id", "") for m in payload.get("data", [])]
    if not any(_model_name_matches(name, OPENROUTER_MODEL) for name in models):
        print(f"[WARN] Model {OPENROUTER_MODEL} not in API list; will verify on first request.")
    print("[OK] OpenRouter connection OK, model ready")
    return True


def choose_llm():
    """
    Return (llm, provider_name).
    Prefer OpenRouter; automatically fall back to Ollama.
    """
    if prefer_openrouter():
        if check_openrouter_ready():
            try:
                test_llm = create_openrouter_llm(max_tokens=4, timeout=30)
                test_llm.invoke("Reply only: OK")
                print("[OK] Using OpenRouter as answer model")
                return create_openrouter_llm(), "openrouter"
            except Exception as e:
                logger.warning("OpenRouter unavailable, falling back to Ollama: %s", e)
                print(f"[WARN] OpenRouter unavailable, falling back to Ollama: {e}")

    try:
        test_llm = create_ollama_llm(num_predict=8, timeout=120)
        test_llm.invoke("ping")
        print(f"[OK] Using Ollama as answer model: {OLLAMA_LLM_MODEL}")
        return create_ollama_llm(), "ollama"
    except Exception as e:
        logger.error("Ollama LLM unavailable: %s", e)
        print(f"[ERROR] Ollama LLM unavailable: {e}")
        print("[TIP] Start Ollama and run: ollama pull llama3")
        return None, None


def warmup_llm() -> None:
    """Send a trivial request so the model is loaded before the first real query."""
    try:
        llm, provider = choose_llm()
        if llm is None:
            return
        llm.invoke("ping")
        print(f"[OK] Model warmup complete ({provider})")
    except Exception as e:
        logger.warning("LLM warmup failed (non-fatal): %s", e)
        print(f"[WARN] Model warmup failed (non-fatal): {e}")

Artigos relacionados