from pydantic import BaseModel, Field, field_validator, EmailStr, ValidationInfo, model_validator
from typing import Optional, List, Literal
from datetime import datetime, date
from enum import Enum



class SortDirection(str, Enum):
    """Enumeration for sort directions"""

    ASC = "asc"
    DESC = "desc"


class DatabaseType(str, Enum):
    """Enumeration for database types"""

    MYSQL = "mysql"
    MONGODB = "mongodb"
    CASSANDRA = "cassandra"

# Prophet Forecasting Schemas

class TimeSeriesDataPoint(BaseModel):
    """Schema for a single time series data point"""

    ds: str = Field(..., description="Date/timestamp in ISO format (YYYY-MM-DD or YYYY-MM-DD HH:MM:SS)")
    y: float = Field(..., description="Numeric value for the time series")

    @field_validator('ds')
    @classmethod
    def validate_date(cls, v):
        """Validate date format"""
        try:
            datetime.fromisoformat(v.replace('Z', '+00:00'))
            return v
        except ValueError:
            raise ValueError('Date must be in ISO format (YYYY-MM-DD or YYYY-MM-DD HH:MM:SS)')


class ProphetConfig(BaseModel):
    """Schema for Prophet model configuration"""

    seasonality_mode: Optional[Literal['additive', 'multiplicative']] = Field(
        'additive', description="Seasonality mode"
    )
    yearly_seasonality: Optional[bool] = Field(True, description="Enable yearly seasonality")
    weekly_seasonality: Optional[bool] = Field(True, description="Enable weekly seasonality")
    daily_seasonality: Optional[bool] = Field(False, description="Enable daily seasonality")
    changepoint_prior_scale: Optional[float] = Field(
        0.05, ge=0.001, le=50, description="Changepoint prior scale"
    )
    seasonality_prior_scale: Optional[float] = Field(
        10.0, ge=0.01, le=100, description="Seasonality prior scale"
    )
    holidays_prior_scale: Optional[float] = Field(
        10.0, ge=0.01, le=100, description="Holidays prior scale"
    )
    changepoint_range: Optional[float] = Field(
        0.8, ge=0.0, le=1.0, description="Changepoint range"
    )
    interval_width: Optional[float] = Field(
        0.80, ge=0.0, le=1.0, description="Confidence interval width (0–1)"
    )


class ProphetTrainRequest(BaseModel):
    """Schema for training a Prophet model"""

    data: List[TimeSeriesDataPoint] = Field(..., min_length=2, description="Time series data points")
    model_id: Optional[str] = Field(None, min_length=1, max_length=100, description="Optional model identifier")
    config: Optional[ProphetConfig] = Field(None, description="Prophet model configuration")


class ProphetForecastRequest(BaseModel):
    """Schema for generating forecasts"""

    periods: Optional[int] = Field(30, ge=1, le=365, description="Number of periods to forecast")
    freq: Optional[Literal['D', 'H', 'W', 'M', 'Q', 'Y']] = Field('D', description="Forecast frequency")
    include_history: Optional[bool] = Field(True, description="Include historical data in response")


class ProphetPlotRequest(BaseModel):
    """Schema for generating forecast plots"""

    periods: Optional[int] = Field(30, ge=1, le=365, description="Number of periods to forecast")
    freq: Optional[Literal['D', 'H', 'W', 'M', 'Q', 'Y']] = Field('D', description="Forecast frequency")
    width: Optional[int] = Field(800, ge=400, le=2000, description="Plot width in pixels")
    height: Optional[int] = Field(600, ge=300, le=1500, description="Plot height in pixels")


class ProphetValidateRequest(BaseModel):
    """Schema for validating time series data"""

    data: List[TimeSeriesDataPoint] = Field(..., min_length=1, description="Time series data points to validate")


# Regression / Classification Schemas

class RegressionTrainRequest(BaseModel):
    """Schema for training a regression or classification model"""
    company_id: int = Field(..., gt=0, description="Company ID to scope training data")
    config: Optional[dict] = Field(None, description="Optional XGBoost hyperparameter overrides")


class CollectionRiskScoreRequest(BaseModel):
    """Schema for scoring client collection risk"""
    company_id: int = Field(..., gt=0, description="Company ID")
    client_ids: Optional[List[int]] = Field(None, description="Specific client IDs to score; omit for all")


class SalesTargetPredictRequest(BaseModel):
    """Schema for predicting sales target attainment"""
    company_id: int = Field(..., gt=0, description="Company ID")
    account_year_id: int = Field(..., gt=0, description="Account year to predict for")
    product_head_ids: Optional[List[int]] = Field(None, description="Filter by product head IDs")
    district_ids: Optional[List[int]] = Field(None, description="Filter by district IDs")