import time
from typing import Dict, Any, Optional
from datetime import datetime
import structlog

logger = structlog.get_logger("app.utils.metrics")


class OperationMetrics:
    def __init__(self, operation_name: str):
        self.operation_name = operation_name
        self.start_time = time.time()
        self.end_time: Optional[float] = None
        self.input_tokens = 0
        self.output_tokens = 0
        self.cost = 0.0
        self.model = ""
        self.cache_hit = False
        self.items_count = 1
        self.error: Optional[str] = None

    def set_tokens(self, input_tokens: int, output_tokens: int):
        self.input_tokens = input_tokens
        self.output_tokens = output_tokens

    def set_cost(self, cost: float):
        self.cost = cost

    def set_model(self, model: str):
        self.model = model

    def set_cache_hit(self, cache_hit: bool):
        self.cache_hit = cache_hit

    def set_items_count(self, count: int):
        self.items_count = count

    def set_error(self, error: str):
        self.error = error

    def finish(self) -> Dict[str, Any]:
        self.end_time = time.time()
        duration = self.end_time - self.start_time
        
        metrics = {
            "operation": self.operation_name,
            "duration": duration,
            "input_tokens": self.input_tokens,
            "output_tokens": self.output_tokens,
            "total_tokens": self.input_tokens + self.output_tokens,
            "cost": self.cost,
            "model": self.model,
            "cache_hit": self.cache_hit,
            "items_count": self.items_count,
            "rate": self.items_count / duration if duration > 0 else 0,
            "timestamp": datetime.now().isoformat(),
            "success": self.error is None
        }
        
        if self.error:
            metrics["error"] = self.error
        
        return metrics

    def log_completion(self):
        metrics = self.finish()
        
        if self.error:
            logger.error(
                "operation_failed",
                **metrics
            )
        else:
            logger.info(
                "operation_completed",
                **metrics
            )


def log_operation_metrics(
    operation: str,
    duration: float,
    input_tokens: int,
    output_tokens: int,
    cost: float,
    model: str,
    items_count: int = 1,
    cache_hit: bool = False,
    error: Optional[str] = None
):
    metrics = {
        "operation": operation,
        "duration": duration,
        "input_tokens": input_tokens,
        "output_tokens": output_tokens,
        "total_tokens": input_tokens + output_tokens,
        "cost": cost,
        "model": model,
        "items_count": items_count,
        "rate": items_count / duration if duration > 0 else 0,
        "cache_hit": cache_hit,
        "timestamp": datetime.now().isoformat(),
        "success": error is None
    }
    
    if error:
        metrics["error"] = error
        logger.error("operation_metrics", **metrics)
    else:
        logger.info("operation_metrics", **metrics)


class BatchMetrics:
    def __init__(self, batch_name: str):
        self.batch_name = batch_name
        self.start_time = time.time()
        self.total_items = 0
        self.successful_items = 0
        self.failed_items = 0
        self.cache_hits = 0
        self.cache_misses = 0
        self.total_input_tokens = 0
        self.total_output_tokens = 0
        self.total_cost = 0.0
        self.errors: list = []

    def add_item_result(self, metrics: Dict[str, Any]):
        self.total_items += 1
        if metrics.get("success", True):
            self.successful_items += 1
        else:
            self.failed_items += 1
            if "error" in metrics:
                self.errors.append(metrics["error"])
        
        if metrics.get("cache_hit", False):
            self.cache_hits += 1
        else:
            self.cache_misses += 1
            
        self.total_input_tokens += metrics.get("input_tokens", 0)
        self.total_output_tokens += metrics.get("output_tokens", 0)
        self.total_cost += metrics.get("cost", 0.0)

    def get_summary(self) -> Dict[str, Any]:
        duration = time.time() - self.start_time
        success_rate = self.successful_items / self.total_items if self.total_items > 0 else 0
        hit_rate = self.cache_hits / (self.cache_hits + self.cache_misses) if (self.cache_hits + self.cache_misses) > 0 else 0
        
        return {
            "batch_name": self.batch_name,
            "total_duration": duration,
            "total_items": self.total_items,
            "successful_items": self.successful_items,
            "failed_items": self.failed_items,
            "success_rate": success_rate,
            "processing_rate": self.successful_items / duration if duration > 0 else 0,
            "cache_hits": self.cache_hits,
            "cache_misses": self.cache_misses,
            "cache_hit_rate": hit_rate,
            "total_input_tokens": self.total_input_tokens,
            "total_output_tokens": self.total_output_tokens,
            "total_cost": self.total_cost,
            "cost_per_item": self.total_cost / self.successful_items if self.successful_items > 0 else 0,
            "errors": self.errors,
            "timestamp": datetime.now().isoformat()
        }

    def log_summary(self):
        summary = self.get_summary()
        logger.info("batch_completed", **summary)