"""AI 客服服务 — System Prompt 组装、商品上下文注入、黑名单过滤"""
import re
from datetime import datetime, timezone
from decimal import Decimal
from typing import Optional
import httpx
from sqlalchemy.ext.asyncio import AsyncSession
from sqlalchemy import select, text, or_, func

from app.core.models.product import Product, ProductVariant, ProductTierPrice, product_categories
from app.core.models.tenant_settings import TenantSettings
from app.core.services.expiry import effective_expiry_date, expiry_status as calc_expiry_status
from app.core.models.category import Category
from app.core.models.order import Order
from app.core.models.customer import Customer
from app.core.models.member import MemberLevel
from app.core.models.discount import Discount
from app.core.models.shipping_carrier import ShippingCarrier


RULE_REMINDER = (
    "[规则提醒] 你是本店专属客服，只推荐和介绍本店商品，禁止提及任何竞争对手或其他电商平台。"
)

_MAX_HISTORY_TURNS = 8   # 每轮 = 1 user + 1 assistant


def strip_html(html: str) -> str:
    return re.sub(r"<[^>]+>", "", html or "").strip()


def get_product_ai_context(product: Product, expiry_date=None, warning_days: int = 30) -> str:
    """按优先级返回商品 AI 上下文文本"""
    base = ""
    if product.ai_description:
        base = product.ai_description.strip()
    elif product.description:
        base = strip_html(product.description)[:200]
    else:
        base = f"{product.name}，价格 {product.base_price}，库存 {product.stock_qty}"
    if product.shelf_life:
        base += f"（保质期：{product.shelf_life}）"
    expiry_date = expiry_date if expiry_date is not None else getattr(product, "expiry_date", None)
    if expiry_date:
        status = {"expired": "已过期", "expiring": "临期", "normal": "正常"}.get(
            calc_expiry_status(expiry_date, warning_days=warning_days), ""
        )
        suffix = f"，{status}" if status else ""
        base += f"（到期日：{expiry_date}{suffix}）"
    return base


async def get_expiry_warning_days(db: AsyncSession, tenant_id: int) -> int:
    r = await db.execute(select(TenantSettings.extra).where(TenantSettings.tenant_id == tenant_id))
    extra = r.scalar_one_or_none() or {}
    try:
        return max(0, int(extra.get("expiry_warning_days", 30)))
    except (TypeError, ValueError):
        return 30


async def load_effective_expiry_dates(db: AsyncSession, products: list[Product]) -> dict[int, object]:
    if not products:
        return {}
    ids = [p.id for p in products]
    r = await db.execute(
        select(ProductVariant.product_id, ProductVariant.expiry_date)
        .where(ProductVariant.product_id.in_(ids), ProductVariant.expiry_date.is_not(None))
    )
    variant_dates: dict[int, list] = {}
    for product_id, expiry_date in r.all():
        variant_dates.setdefault(product_id, []).append(expiry_date)
    return {
        p.id: effective_expiry_date(getattr(p, "expiry_date", None), variant_dates.get(p.id, []))
        for p in products
    }


async def get_member_level(db: AsyncSession, customer: Optional[Customer]) -> Optional[MemberLevel]:
    if not customer or not customer.member_level_id:
        return None
    r = await db.execute(select(MemberLevel).where(MemberLevel.id == customer.member_level_id))
    return r.scalar_one_or_none()


def _effective_price(
    product: Product,
    member_level: Optional[MemberLevel],
    tier_price: Optional[Decimal] = None,
) -> Decimal:
    """按优先级计算商品有效价格（与 pricing.py 逻辑一致）"""
    base = product.base_price
    if tier_price is not None:
        return tier_price
    if product.member_price is not None and member_level:
        effective = product.member_price
        if member_level.discount_rate < Decimal("1.0000"):
            discounted = base * member_level.discount_rate
            if discounted < effective:
                effective = discounted
        return effective
    if member_level and member_level.discount_rate < Decimal("1.0000"):
        return (base * member_level.discount_rate).quantize(Decimal("0.01"))
    return base


def _price_str(base: Decimal, effective: Decimal, member_level_name: str, sym: str = "") -> str:
    if effective < base:
        return f"{sym}{effective}（原价{sym}{base}，{member_level_name}专享）"
    return f"{sym}{base}"


async def _load_products_with_price(
    db: AsyncSession,
    products: list[Product],
    member_level: Optional[MemberLevel],
) -> list[tuple[Product, Decimal]]:
    """返回 (product, effective_price) 列表"""
    tier_map: dict[int, Decimal] = {}
    if member_level and products:
        ids = [p.id for p in products]
        tp_r = await db.execute(
            select(ProductTierPrice.product_id, ProductTierPrice.price).where(
                ProductTierPrice.product_id.in_(ids),
                ProductTierPrice.member_level_id == member_level.id,
                ProductTierPrice.variant_id.is_(None),
            )
        )
        tier_map = {row[0]: row[1] for row in tp_r.all()}
    return [(p, _effective_price(p, member_level, tier_map.get(p.id))) for p in products]


def _sanitize_fulltext_kw(kw: str) -> str:
    """去掉 MySQL FULLTEXT BOOLEAN MODE 的特殊运算符，避免语法错误"""
    for ch in ('+', '-', '>', '<', '(', ')', '~', '*', '"', '@'):
        kw = kw.replace(ch, ' ')
    return kw.strip()


async def fetch_product_context(
    db: AsyncSession,
    keyword: str,
    tenant_id: int,
    member_level: Optional[MemberLevel] = None,
) -> str:
    """关键词 FULLTEXT 检索相关商品，按会员等级返回有效价格"""
    from app.core.models.currency import Currency
    _sym_r = await db.execute(
        select(Currency.symbol).where(Currency.tenant_id == tenant_id, Currency.is_default == 1).limit(1)
    )
    _sym = _sym_r.scalar_one_or_none() or ""
    level_name = member_level.name if member_level else ""
    products: list[Product] = []
    used_fulltext = False

    _more_hint = ""   # 当结果超过12条时，向 AI 注入"搜索更多"指令

    if keyword and keyword.strip():
        kw = keyword.strip()

        # ① LIKE 精确匹配（主要检索，2字以上 token，避免单字误匹配）
        # 先用较长的 token 精确搜，精度高；搜不到再用单字兜底。
        try:
            _STOPCHARS = set('有没吗你我他她它的了是在不都也就和与或但很为什么怎么如何'
                             '哪里谁这那啊呢嗯哦哈想买请问一个些要啥多少价格推荐介绍')

            def _build_like_conds(tokens):
                conds = []
                for token in tokens:
                    conds.append(Product.name.ilike(f"%{token}%"))
                    conds.append(Product.name_en.ilike(f"%{token}%"))
                    conds.append(Product.ai_description.ilike(f"%{token}%"))
                    conds.append(Product.ai_description_en.ilike(f"%{token}%"))
                return conds

            seen: set[str] = set()
            long_tokens: list[str] = []   # 2字及以上
            short_tokens: list[str] = []  # 1字 fallback
            for length in range(min(4, len(kw)), 0, -1):
                for i in range(len(kw) - length + 1):
                    token = kw[i:i + length]
                    if token in seen:
                        continue
                    if all(c in _STOPCHARS for c in token):
                        continue
                    seen.add(token)
                    if length >= 2:
                        long_tokens.append(token)
                    else:
                        short_tokens.append(token)
                if len(long_tokens) + len(short_tokens) >= 20:
                    break

            # 提取最佳搜索关键词：优先取最短的纯非停用字 token（如"酱油"而非"有酱油"）
            _search_kw = ""
            for t in sorted(long_tokens, key=len):
                if not any(c in _STOPCHARS for c in t):
                    _search_kw = t
                    break
            if not _search_kw and long_tokens:
                _search_kw = min(long_tokens, key=len)

            # Phase 1: 用 2字以上 token 搜，limit 12，同时 COUNT 总数
            if long_tokens:
                conds = _build_like_conds(long_tokens)
                base_where = [
                    Product.tenant_id == tenant_id,
                    Product.status == "active",
                    or_(*conds),
                ]
                # 先 COUNT 总数
                from sqlalchemy import func as _func
                cnt_r = await db.execute(select(_func.count()).select_from(Product).where(*base_where))
                _total_count = cnt_r.scalar() or 0

                r = await db.execute(
                    select(Product).where(*base_where)
                    .order_by(Product.sales_count.desc()).limit(12)
                )
                products = r.scalars().all()
                if products and _total_count > 12 and _search_kw:
                    _more_hint = "\n【搜索更多】本次共找到 %d 个相关商品，以上展示了前 12 个。请在回复末尾加上搜索跳转标签：[SEARCH|%s|查看全部 %d 个%s商品]" % (_total_count, _search_kw, _total_count, _search_kw)
                else:
                    _more_hint = ""

            # Phase 2: 没结果时才用单字 token 兜底
            if not products and short_tokens:
                _more_hint = ""
                conds = _build_like_conds(short_tokens)
                r = await db.execute(
                    select(Product).where(
                        Product.tenant_id == tenant_id,
                        Product.status == "active",
                        or_(*conds),
                    ).order_by(Product.sales_count.desc()).limit(12)
                )
                products = r.scalars().all()

            if products:
                used_fulltext = True
        except Exception:
            _more_hint = ""
            pass

        # ② MySQL FULLTEXT（LIKE 未命中时的补充，对纯英文关键词有效）
        if not products:
            try:
                safe_kw = _sanitize_fulltext_kw(kw)
                if safe_kw:
                    q = text(
                        "SELECT id FROM products "
                        "WHERE tenant_id = :tid AND status = 'active' "
                        "AND MATCH(name, description, ai_description) "
                        "AGAINST (:kw IN BOOLEAN MODE) "
                        "LIMIT 12"
                    )
                    result = await db.execute(q, {"tid": tenant_id, "kw": safe_kw})
                    ids = [row[0] for row in result.all()]
                    if ids:
                        r2 = await db.execute(select(Product).where(Product.id.in_(ids)))
                        products = r2.scalars().all()
                        used_fulltext = True
            except Exception:
                pass

    # 如果 LIKE/FULLTEXT 均未命中，尝试按分类名称匹配（用户消息包含分类名，含英文分类名）
    if not products and keyword and keyword.strip():
        _more_hint = ""
        try:
            kw = keyword.strip()
            from sqlalchemy import func as sqlfunc
            cats_r = await db.execute(
                select(Category.id).where(
                    Category.tenant_id == tenant_id,
                    or_(
                        sqlfunc.instr(kw, Category.name) > 0,
                        sqlfunc.instr(kw, Category.name_en) > 0,
                    ),
                )
            )
            cat_ids = [row[0] for row in cats_r.all()]
            if cat_ids:
                cat_prod_r = await db.execute(
                    select(Product)
                    .join(product_categories, Product.id == product_categories.c.product_id)
                    .where(
                        Product.tenant_id == tenant_id,
                        Product.status == "active",
                        product_categories.c.category_id.in_(cat_ids),
                    )
                    .order_by(Product.sales_count.desc())
                    .limit(12)
                )
                products = cat_prod_r.scalars().all()
                if products:
                    used_fulltext = True
        except Exception:
            pass

    if not products:
        _more_hint = ""
        r = await db.execute(
            select(Product).where(
                Product.tenant_id == tenant_id,
                Product.status == "active",
            ).order_by(Product.sales_count.desc()).limit(5)
        )
        products = r.scalars().all()

    if not products:
        return ""

    pairs = await _load_products_with_price(db, products, member_level)
    warning_days = await get_expiry_warning_days(db, tenant_id)
    expiry_dates = await load_effective_expiry_dates(db, products)
    lines = [
        f"- 《{p.name}|{p.slug}|{p.id}》（{_price_str(p.base_price, price, level_name, _sym)}，库存{p.stock_qty}）：{get_product_ai_context(p, expiry_dates.get(p.id), warning_days)}"
        for p, price in pairs
    ]
    prefix = "【相关商品】" if used_fulltext else "【推荐商品】"
    result_text = prefix + "\n" + "\n".join(lines)
    if _more_hint:
        result_text += _more_hint
    return result_text


async def fetch_order_context(db: AsyncSession, customer: Optional[Customer]) -> str:
    """返回最近 3 单订单摘要；未登录时返回引导登录提示"""
    if customer is None:
        return "【订单查询】用户当前未登录，无法查询订单。若用户询问订单，请引导其前往登录页面完成登录后再查询。"
    q = select(Order).where(
        Order.customer_id == customer.id,
    ).order_by(Order.created_at.desc()).limit(3)
    result = await db.execute(q)
    orders = result.scalars().all()
    if not orders:
        return "【订单查询】该用户暂无订单记录。"
    lines = []
    for o in orders:
        tracking = ""
        if o.tracking_no:
            tracking = f"，物流单号：{o.tracking_no}"
        shipments = (o.extra_attributes or {}).get("shipments", [])
        if shipments:
            carriers = "；".join(
                f"{s.get('carrier', '')} {s.get('tracking_no', '')}" for s in shipments
            )
            tracking = f"，发货信息：{carriers}"
        lines.append(f"- 订单 {o.order_no}：状态={o.status}，金额=¥{o.grand_total}{tracking}")
    return "【我的订单】\n" + "\n".join(lines)


async def fetch_promotion_context(db: AsyncSession, tenant_id: int) -> str:
    """返回当前有效的促销活动摘要"""
    now = datetime.now(timezone.utc)
    r = await db.execute(
        select(Discount).where(
            Discount.tenant_id == tenant_id,
            Discount.is_active == 1,
            or_(Discount.start_at.is_(None), Discount.start_at <= now),
            or_(Discount.end_at.is_(None), Discount.end_at >= now),
        ).order_by(Discount.priority.desc()).limit(8)
    )
    discounts = r.scalars().all()
    if not discounts:
        return ""

    _type_desc = {
        "percentage":        lambda d: f"享{round(float(d.value))}%折扣",
        "fixed":             lambda d: f"减¥{d.value}",
        "free_shipping":     lambda d: "免运费",
        "bogo":              lambda d: f"买赠（赠{d.discount_qty or 1}件）",
        "points_multiplier": lambda d: f"{d.points_multiplier}倍积分",
    }
    lines = []
    for d in discounts:
        desc_fn = _type_desc.get(d.type, lambda d: d.type)
        desc = desc_fn(d)
        cond = f"满¥{d.min_order_amount}" if d.min_order_amount else ""
        if d.code:
            lines.append(f"- 【优惠码】{d.name}：输入码「{d.code}」{cond}{desc}")
        else:
            lines.append(f"- 【自动优惠】{d.name}：{cond}{desc}")
    return "【当前优惠活动】\n" + "\n".join(lines)


def build_system_prompt(
    shop_name: str,
    chat_name: str,
    faq: str,
    product_ctx: str,
    order_ctx: str,
    promo_ctx: str = "",
    member_level_name: str = "",
    page_ctx: str = "",
    tracking_ctx: str = "",
) -> str:
    parts = [
        f"你是【{shop_name}】的专属客服助手，名叫【{chat_name}】。",
        "严格规则：",
        "1. 只介绍和推荐本店商品，禁止提及任何竞争对手品牌或其他电商平台（如京东、淘宝、拼多多、亚马逊等）。",
        "2. 商品推荐规则：下方【相关商品】是系统根据用户问题检索出的结果。"
        "若列表中有用户询问的商品（完全或部分匹配），直接介绍它，不要说'暂无'。"
        "若列表中没有完全匹配的商品但有相关商品，直接介绍相关商品，说'为您找到以下相关商品'，不要说'暂无XX'。"
        "只有列表中完全没有任何相关商品时，才说本店暂无该类商品，并询问用户其他需求。",
        "3. 不回答与购物完全无关的话题（天气、政治、技术问题等），礼貌引导回购物场景。",
        "4. 不得捏造商品信息，如不确定请说'请联系人工客服确认'。",
        "5. 回复简洁友好，使用中文。报价时必须使用下方【相关商品】或【推荐商品】中的实际价格，不得使用其他数字。",
        "6. 提及具体商品时，必须原样保留并输出 《商品名|slug|id》 格式标签（如 《厨邦酱油|chubang-soy-sauce|42》），不得修改、拆分或省略，用户可通过该标签直接点击跳转和加购。",
        '7. 若商品上下文中包含【搜索更多】指令（含 [SEARCH|keyword|text] 格式），必须在回复末尾原样输出该标签，例如：[SEARCH|酱油|查看全部 9 个酱油商品]。不得修改关键词或文字内容。',
    ]
    if member_level_name:
        parts.append(f"\n【当前用户】会员等级：{member_level_name}，报价时请使用对应会员专享价（已在商品列表中标注）。")
    if faq:
        parts.append(f"\n【商家FAQ / 政策】\n{faq}")
    if promo_ctx:
        parts.append(f"\n{promo_ctx}")
    if page_ctx:
        parts.append(f"\n{page_ctx}")
    if product_ctx:
        parts.append(f"\n{product_ctx}")
    if order_ctx:
        parts.append(f"\n{order_ctx}")
    if tracking_ctx:
        parts.append(f"\n{tracking_ctx}")
    return "\n".join(parts)


def apply_sliding_window(messages: list[dict]) -> tuple[list[dict], bool]:
    """
    保留最近 _MAX_HISTORY_TURNS 轮对话历史。
    返回 (裁剪后的历史, 是否发生了截断)。
    """
    max_msgs = _MAX_HISTORY_TURNS * 2
    if len(messages) <= max_msgs:
        return messages, False
    return messages[-max_msgs:], True


def filter_blacklist(text: str, blacklist: str) -> str:
    """检测回复中是否出现竞品关键词，命中则替换为 ***"""
    if not blacklist:
        return text
    for kw in (k.strip() for k in blacklist.split(",") if k.strip()):
        if kw in text:
            text = text.replace(kw, "***")
            break
    return text


# 触发快递查询的关键词（中英文）
_TRACKING_KEYWORDS = {
    '快递', '物流', '发货', '运单', '到了', '到哪', '到没', '寄出',
    '派送', '配送', '邮寄', '包裹', '签收', 'tracking', 'delivery',
    'shipped', 'parcel', 'courier', 'express',
}


def _is_tracking_query(text: str) -> bool:
    low = text.lower()
    return any(kw in low for kw in _TRACKING_KEYWORDS)


def _strip_html_text(html: str) -> str:
    """去掉 HTML 标签，压缩空白，返回纯文本（限 1500 字）
    优先提取 class 含 query-content / timeline / track 的区块，避免页头导航占用篇幅。
    """
    # 尝试提取物流内容区块：找到起始位置后截取 6000 字节
    for keyword in ('query-content', 'timeline', 'track-result'):
        idx = html.lower().find(keyword)
        if idx != -1:
            html = html[max(0, idx - 50): idx + 6000]
            break
    text = re.sub(r'<[^>]+>', ' ', html)
    text = re.sub(r'\s+', ' ', text).strip()
    return text[:2000]


async def fetch_tracking_context(
    db: AsyncSession,
    customer: Optional[Customer],
    tenant_id: int,
    last_user_msg: str,
) -> str:
    """
    检测用户消息是否询问快递，若是则查近期有运单号的订单，
    抓取 tracking URL 页面内容，返回物流上下文字符串。
    """
    if not customer:
        return ""
    if not _is_tracking_query(last_user_msg):
        return ""

    r = await db.execute(
        select(Order).where(
            Order.customer_id == customer.id,
            Order.tenant_id == tenant_id,
            Order.tracking_no.is_not(None),
            Order.tracking_no != "",
        ).order_by(Order.created_at.desc()).limit(5)
    )
    orders = r.scalars().all()
    if not orders:
        return ""

    lines: list[str] = ["[物流信息]"]

    for order in orders:
        tracking_no = order.tracking_no or ""
        carrier_name = order.carrier or ""
        tracking_url: str | None = None

        if carrier_name and tracking_no:
            rc = await db.execute(
                select(ShippingCarrier).where(
                    ShippingCarrier.tenant_id == tenant_id,
                    ShippingCarrier.is_active == 1,
                    or_(
                        func.lower(ShippingCarrier.name) == carrier_name.lower(),
                        func.lower(ShippingCarrier.code) == carrier_name.lower(),
                    ),
                )
            )
            carrier_obj = rc.scalar_one_or_none()
            if carrier_obj and carrier_obj.tracking_url_template:
                tracking_url = carrier_obj.tracking_url_template.replace("{tracking_no}", tracking_no)

        line = f"订单 {order.order_no}，运单号 {tracking_no}"
        if carrier_name:
            line += f"（{carrier_name}）"

        if tracking_url:
            try:
                async with httpx.AsyncClient(timeout=25) as client:
                    resp = await client.get(tracking_url, follow_redirects=True)
                    if resp.status_code == 200:
                        page_text = _strip_html_text(resp.text)
                        if page_text:
                            line += f"，物流详情：{page_text}"
            except Exception:
                pass  # 抓取失败则只返回运单号，不中断

        lines.append(line)

    return "\n".join(lines)
