Explorer
/opt/struktur/lead-engine/llm/router.py
← Zurück ↓ Download
"""
Router für LLM-Enrichment — ausschließlich OpenAI.
3-Level Routing: gpt-5-nano → gpt-5-mini → gpt-5.2

Kein Anthropic. Kein externer Fallback.
"""

import os
import logging
from typing import Dict, Any, Optional, Tuple

from .validator import validate, get_estimated_fields
from .logger import enrichment_logger
from .openai_provider import OpenAIProvider

logger = logging.getLogger(__name__)


# Cost Guard Parameter (von Dottore)
MIN_SAMPLE         = 50    # Erst ab N Requests aktiv
WARN_THRESHOLD     = 0.15  # >15% → Warnung, kein Eingriff
THROTTLE_THRESHOLD = 0.25  # >25% → Level 3 nur bei kritischen Fehlern
STOP_THRESHOLD     = 0.35  # >35% → Hard Stop


class EnrichmentRouter:
    """
    Router mit 3 OpenAI-Leveln und 4-Phasen Cost Guard.

    HARTE REGELN:
    1. Immer mit Level 1 (günstig) starten
    2. Jeder Provider genau 1 Versuch
    3. Validierungsfehler → Eskalation (außer bei Throttle)
    4. Alle Level fail → None (kein save)
    5. Cost Guard: 4 Phasen (normal → warning → throttled → stop)
    """

    def __init__(self):
        self.levels = [
            {'level': 1, 'model_env': 'OPENAI_MODEL_LOW',  'provider': None},
            {'level': 2, 'model_env': 'OPENAI_MODEL_MID',  'provider': None},
            {'level': 3, 'model_env': 'OPENAI_MODEL_HIGH', 'provider': None},
        ]

        # Abbruchlogik
        self.consecutive_fails = 0
        self.max_consecutive_fails = 3

        # Cost Guard Zähler
        self.total_counter  = 0
        self.level3_counter = 0

        # Performance (Downpriorisierung)
        self.downprioritized = set()

        logger.info("EnrichmentRouter (OpenAI-only) initialized: nano → mini → gpt-5.2")

    def _get_provider(self, level_config: Dict) -> Optional[OpenAIProvider]:
        """Lazy-load OpenAI Provider."""
        if level_config['provider'] is None:
            model = os.getenv(level_config['model_env'])
            if not model:
                logger.error(f"Env var {level_config['model_env']} nicht gesetzt")
                return None
            try:
                level_config['provider'] = OpenAIProvider(model)
                logger.debug(f"Provider L{level_config['level']} geladen: {model}")
            except Exception as e:
                logger.error(f"Provider L{level_config['level']} Fehler: {e}")
                return None
        return level_config['provider']

    def _get_level3_ratio(self) -> float:
        """Berechnet aktuellen Level-3 Anteil."""
        if self.total_counter == 0:
            return 0.0
        return self.level3_counter / self.total_counter

    def _get_cost_guard_state(self) -> str:
        """
        Bestimmt den aktuellen Cost-Guard-Zustand.

        Returns:
            'normal' | 'warning' | 'throttled' | 'stop'
        """
        if self.total_counter < MIN_SAMPLE:
            return 'normal'

        ratio = self._get_level3_ratio()

        if ratio > STOP_THRESHOLD:
            return 'stop'
        if ratio > THROTTLE_THRESHOLD:
            return 'throttled'
        if ratio > WARN_THRESHOLD:
            return 'warning'
        return 'normal'

    def _check_cost_guard(self):
        """
        4-Phasen Cost Guard nach Dottores Spezifikation.

        Phase 1 (< MIN_SAMPLE):   Kein Eingriff, nur Logging
        Phase 2 (> WARN):         Warnung, kein Eingriff
        Phase 3 (> THROTTLE):     Level 3 eingeschränkt (→ nur kritische Fehler)
        Phase 4 (> STOP):         Hard Stop, raise Exception
        """
        state = self._get_cost_guard_state()
        ratio = self._get_level3_ratio()

        if state == 'normal':
            return  # Kein Eingriff

        if state == 'warning':
            logger.warning(
                f"LEVEL3_HIGH_USAGE: {self.level3_counter}/{self.total_counter} "
                f"({ratio*100:.1f}%) — Warnschwelle {WARN_THRESHOLD*100:.0f}% überschritten. "
                f"Kein Eingriff."
            )
            return

        if state == 'throttled':
            logger.warning(
                f"COST_GUARD_THROTTLED: Level-3 Anteil {ratio*100:.1f}% > "
                f"{THROTTLE_THRESHOLD*100:.0f}%. Level 3 nur noch bei kritischen Fehlern."
            )
            return  # Eingriff erfolgt in run_enrichment

        if state == 'stop':
            raise Exception(
                f"COST_GUARD_STOP: Level-3 Anteil {ratio*100:.1f}% > "
                f"{STOP_THRESHOLD*100:.0f}%. Prozess gestoppt."
            )

    def run_enrichment(self, prompt: str, contact_id: int,
                       max_budget: Optional[int] = None) -> Tuple[bool, Optional[Dict]]:
        """
        Führt Enrichment durch: Level 1 → 2 → 3.

        Args:
            max_budget: Token-Budget pro Kontakt (Override für .env MAX_TOKENS_PER_CONTACT)

        Returns:
            (True, data) bei Erfolg
            (False, None) wenn alle Level scheitern oder Budget überschritten
        """
        # Token-Budget aus ENV laden (Override möglich für Multi-Source)
        if max_budget is None:
            max_budget = int(os.getenv('MAX_TOKENS_PER_CONTACT', '25000'))

        budget_used = 0
        tokens_per_level = {}

        # Abbruch nach 3 konsekutiven Totalausfällen
        if self.consecutive_fails >= self.max_consecutive_fails:
            logger.error(
                f"[{contact_id}] PROCESS_STOPPED: "
                f"{self.max_consecutive_fails} konsekutive Totalausfälle"
            )
            return False, None

        # Cost Guard prüfen
        try:
            self._check_cost_guard()
        except Exception as e:
            logger.critical(f"[{contact_id}] Cost Guard ausgelöst: {e}")
            return False, None

        self.total_counter += 1
        cost_guard_state = self._get_cost_guard_state()

        # Tracking: letzter Fehlertyp pro Level (für Throttle-Entscheidung)
        last_fail_reason = None

        for level_config in self.levels:
            level_num = level_config['level']

            # BUDGET CHECK: vor jedem API-Call
            if budget_used >= max_budget:
                logger.warning(
                    f"[{contact_id}] BUDGET_EXCEEDED: {budget_used}/{max_budget} Tokens "
                    f"nach L{level_num - 1} → needs_manual"
                )
                enrichment_logger.log_contact_budget(
                    contact_id=contact_id,
                    budget_used=budget_used,
                    max_budget=max_budget,
                    tokens_per_level=tokens_per_level,
                    exceeded=True,
                )
                return False, None

            # THROTTLE: Level 3 nur bei kritischen Fehlern (nicht bei Validierung)
            if level_num == 3 and cost_guard_state == 'throttled':
                if last_fail_reason and last_fail_reason.startswith('MISSING_CRITICAL'):
                    logger.info(
                        f"[{contact_id}] Throttled aber CRITICAL fehlt → Level 3 erlaubt"
                    )
                elif last_fail_reason and 'API_ERROR' in last_fail_reason:
                    logger.info(
                        f"[{contact_id}] Throttled aber API-Fehler → Level 3 erlaubt"
                    )
                else:
                    logger.info(
                        f"[{contact_id}] THROTTLE: Level 3 übersprungen "
                        f"(Grund: {last_fail_reason})"
                    )
                    break

            provider = self._get_provider(level_config)
            if provider is None:
                continue

            model_name = provider.model

            # Downpriorisierte Modelle überspringen (außer Level 3)
            if model_name in self.downprioritized and level_num < 3:
                logger.debug(f"[{contact_id}] Skip downpriorisiert: {model_name}")
                continue

            try:
                # 1. API Call + JSON Parse (in Provider)
                data = provider.extract(prompt)
                tokens = getattr(provider, 'tokens_used', 0)

                # Budget akkumulieren
                budget_used += tokens
                tokens_per_level[f"L{level_num}"] = tokens

                # Level-3 Zähler
                if level_num == 3:
                    self.level3_counter += 1

                # 2. Validierung
                valid, reason = validate(data)
                estimated = get_estimated_fields(data)
                last_fail_reason = None if valid else reason

                # 3. Logging
                enrichment_logger.log_attempt(
                    contact_id=contact_id,
                    model_used=model_name,
                    level=level_num,
                    success=valid,
                    validation_result=reason,
                    tokens_used=tokens,
                    estimated_fields=estimated,
                )

                if valid:
                    self.consecutive_fails = 0
                    enrichment_logger.log_contact_budget(
                        contact_id=contact_id,
                        budget_used=budget_used,
                        max_budget=max_budget,
                        tokens_per_level=tokens_per_level,
                        exceeded=False,
                    )
                    logger.info(
                        f"[{contact_id}] SUCCESS mit {model_name} (L{level_num}) | "
                        f"Budget: {budget_used}/{max_budget} Tokens"
                    )
                    return True, data

                logger.debug(
                    f"[{contact_id}] Validierung fehlgeschlagen auf {model_name}: "
                    f"{reason} → Eskalation"
                )

            except Exception as e:
                tokens = getattr(provider, 'tokens_used', 0)
                budget_used += tokens
                tokens_per_level[f"L{level_num}"] = tokens
                last_fail_reason = f"API_ERROR: {str(e)[:50]}"
                enrichment_logger.log_attempt(
                    contact_id=contact_id,
                    model_used=model_name,
                    level=level_num,
                    success=False,
                    validation_result='API_ERROR',
                    tokens_used=tokens,
                    error_msg=str(e),
                )
                logger.warning(
                    f"[{contact_id}] API-Fehler auf {model_name}: {e} → Eskalation"
                )

        # Alle Level fehlgeschlagen
        self.consecutive_fails += 1
        enrichment_logger.log_contact_budget(
            contact_id=contact_id,
            budget_used=budget_used,
            max_budget=max_budget,
            tokens_per_level=tokens_per_level,
            exceeded=False,
        )
        logger.error(
            f"[{contact_id}] ALLE LEVEL FEHLGESCHLAGEN. "
            f"Konsekutiv: {self.consecutive_fails}/{self.max_consecutive_fails}"
        )

        if self.consecutive_fails >= self.max_consecutive_fails:
            logger.critical(
                f"Prozess wird gestoppt nach {self.max_consecutive_fails} "
                f"konsekutiven Totalausfällen."
            )

        return False, None

    def check_model_performance(self, threshold: int = 50):
        """
        Nach N Requests: Modelle mit <70% Erfolg downpriorisieren.
        Bleibt als Fallback aktiv — wird nur nicht mehr als Level 1 genutzt.
        """
        for model, stats in enrichment_logger.stats.items():
            if stats['attempts'] >= threshold:
                rate = stats['successes'] / stats['attempts']
                if rate < 0.70 and model not in self.downprioritized:
                    self.downprioritized.add(model)
                    logger.warning(
                        f"Model {model}: {rate*100:.1f}% < 70% nach "
                        f"{stats['attempts']} Requests → downpriorisiert"
                    )

    def get_stats(self) -> Dict:
        ratio = self._get_level3_ratio()
        return {
            'consecutive_fails': self.consecutive_fails,
            'total_requests': self.total_counter,
            'level3_requests': self.level3_counter,
            'level3_ratio': round(ratio, 3),
            'cost_guard_state': self._get_cost_guard_state(),
            'downprioritized': list(self.downprioritized),
            'model_stats': enrichment_logger.stats,
        }


# Global Instance
router = EnrichmentRouter()


def run_enrichment(prompt: str, contact_id: int) -> Tuple[bool, Optional[Dict]]:
    return router.run_enrichment(prompt, contact_id)


def check_performance(threshold: int = 50):
    router.check_model_performance(threshold)


def print_stats():
    enrichment_logger.print_stats()
    stats = router.get_stats()
    state = stats['cost_guard_state']
    state_icons = {
        'normal':    '[OK]',
        'warning':   '[WARN]',
        'throttled': '[THROTTLE]',
        'stop':      '[STOP]',
    }
    print(f"\nRouter Status:")
    print(f"  Konsekutive Fehler:  {stats['consecutive_fails']}/3")
    print(f"  Total Requests:      {stats['total_requests']}")
    print(f"  Level-3 Anteil:      {stats['level3_requests']}/{stats['total_requests']} "
          f"({stats['level3_ratio']*100:.1f}%)")
    print(f"  Cost Guard:          {state_icons.get(state, state)} {state.upper()} "
          f"[Schwellen: {WARN_THRESHOLD*100:.0f}% / {THROTTLE_THRESHOLD*100:.0f}% / {STOP_THRESHOLD*100:.0f}%]")
    print(f"  Downpriorisiert:     {stats['downprioritized']}")