"""Final SOW AI helpers: inclusive JD rewrite, next-question, TTS, scenarios, lock validate,
follow-up logic, merge custom questions.
"""

from uuid import uuid4

from app.schemas.common import JobStatus
from app.schemas.questions import QuestionType
from app.schemas.helpers import (
    FollowUpLogicRequest,
    FollowUpLogicResponse,
    FollowUpRule,
    InterviewValidateLockRequest,
    InterviewValidateLockResponse,
    JdInclusivityRewriteRequest,
    JdInclusivityRewriteResponse,
    MergeCustomQuestionsRequest,
    MergeCustomQuestionsResponse,
    NextQuestionRequest,
    NextQuestionResponse,
    ScenarioGenerateRequest,
    ScenarioGenerateResponse,
    ScenarioGradeRequest,
    ScenarioGradeResponse,
    TtsSynthesizeRequest,
    TtsSynthesizeResponse,
)
from app.services.ai_generate import generate_json
from app.services.questions import _local_questions, _parse_questions


class JdInclusivityService:
    async def rewrite(
        self, payload: JdInclusivityRewriteRequest
    ) -> JdInclusivityRewriteResponse:
        system = (
            "Rewrite JD to be inclusive and bias-free. Return JSON: "
            "{\"rewritten_markdown\",\"removed_bias_phrases\":[],\"notes\":[]}"
        )
        user = f"Title: {payload.role_title}\nJD:\n{payload.job_description}"

        def local() -> dict:
            risky = ["young blood", "native speaker", "guys", "rockstar ninja"]
            found = [p for p in risky if p in payload.job_description.lower()]
            text = payload.job_description
            for p in found:
                text = text.replace(p, "skilled professional")
                text = text.replace(p.title(), "Skilled professional")
            return {
                "rewritten_markdown": text,
                "removed_bias_phrases": found,
                "notes": ["Prefer skill/outcome language over demographic proxies"],
            }

        data, stub, provider = await generate_json(
            system=system, user=user, local_factory=local
        )
        return JdInclusivityRewriteResponse(
            correlation_id=payload.correlation_id,
            status=JobStatus.SUCCEEDED,
            stub=stub,
            message=f"Inclusivity rewrite via {provider}",
            original=payload.job_description,
            rewritten_markdown=data.get("rewritten_markdown"),
            removed_bias_phrases=list(data.get("removed_bias_phrases") or []),
            notes=list(data.get("notes") or []),
        )


class NextQuestionService:
    async def choose(self, payload: NextQuestionRequest) -> NextQuestionResponse:
        system = (
            "Pick the best next interview question. Return JSON: "
            "{\"question_id\",\"question_text\",\"rationale\",\"done\":bool}"
        )
        user = (
            f"Role: {payload.role_title}\nJD: {payload.job_description}\n"
            f"Planned: {payload.planned_questions}\nAsked: {payload.asked_question_ids}\n"
            f"Last score: {payload.last_score}\nTranscript: {payload.transcript[-5:]}"
        )

        def local() -> dict:
            remaining = [
                q
                for i, q in enumerate(payload.planned_questions)
                if f"q-{i + 1}" not in payload.asked_question_ids
                and q not in payload.asked_question_ids
            ]
            if not remaining and not payload.planned_questions:
                return {
                    "question_id": "nq-1",
                    "question_text": (
                        f"What recent impact did you deliver as a {payload.role_title}?"
                    ),
                    "rationale": "No planned list — generate adaptive probe",
                    "done": False,
                }
            if not remaining:
                return {
                    "question_id": None,
                    "question_text": None,
                    "rationale": "All planned questions asked",
                    "done": True,
                }
            idx = 0
            if payload.last_score is not None and payload.last_score < 0.6 and len(remaining) > 1:
                idx = 0
            qtext = remaining[idx]
            qid = f"q-{payload.planned_questions.index(qtext) + 1}"
            return {
                "question_id": qid,
                "question_text": qtext,
                "rationale": "Next unanswered planned question",
                "done": False,
            }

        data, stub, provider = await generate_json(
            system=system, user=user, local_factory=local
        )
        return NextQuestionResponse(
            correlation_id=payload.correlation_id,
            interview_id=payload.interview_id,
            status=JobStatus.SUCCEEDED,
            stub=stub,
            message=f"Next question via {provider}",
            question_id=data.get("question_id"),
            question_text=data.get("question_text"),
            rationale=data.get("rationale"),
            done=bool(data.get("done")),
        )


class TtsService:
    async def synthesize(self, payload: TtsSynthesizeRequest) -> TtsSynthesizeResponse:
        return TtsSynthesizeResponse(
            correlation_id=payload.correlation_id,
            status=JobStatus.SUCCEEDED,
            stub=True,
            message="TTS local stub — live audio provider not wired yet",
            audio_base64=None,
            mime_type="audio/mpeg",
            voice_id=payload.voice_id or payload.voice_name or "default",
            provider="local",
        )


class ScenariosService:
    async def generate(
        self, payload: ScenarioGenerateRequest
    ) -> ScenarioGenerateResponse:
        system = "Generate role-based scenario questions. Return JSON: {\"scenarios\":[...]}"
        user = (
            f"Role: {payload.role_title}\nJD: {payload.job_description}\n"
            f"Difficulty: {payload.difficulty.value}\nCount: {payload.count}"
        )

        def local() -> dict:
            return {
                "scenarios": _local_questions(
                    count=payload.count,
                    qtype=QuestionType.LONG_ANSWER,
                    difficulty=payload.difficulty,
                    topic=f"scenario / {payload.role_title}",
                )
            }

        data, stub, provider = await generate_json(
            system=system, user=user, local_factory=local
        )
        scenarios = _parse_questions(
            data.get("scenarios", []), QuestionType.LONG_ANSWER, payload.difficulty
        )
        return ScenarioGenerateResponse(
            correlation_id=payload.correlation_id,
            scenario_set_id=str(uuid4()),
            status=JobStatus.SUCCEEDED,
            stub=stub,
            message=f"Scenarios via {provider}",
            scenarios=scenarios,
        )

    async def grade(self, payload: ScenarioGradeRequest) -> ScenarioGradeResponse:
        system = (
            "Grade scenario answers. Return JSON: "
            "{\"total_score\",\"max_total_score\":100,\"passed\",\"feedback_markdown\"}"
        )
        user = f"Role: {payload.role_title}\nAnswers: {payload.answers}"

        def local() -> dict:
            n = len(payload.answers)
            score = min(95.0, 50.0 + 10.0 * n)
            return {
                "total_score": score,
                "max_total_score": 100.0,
                "passed": score >= 70,
                "feedback_markdown": "Scenarios show structured thinking; add trade-off analysis.",
            }

        data, stub, provider = await generate_json(
            system=system, user=user, local_factory=local
        )
        return ScenarioGradeResponse(
            correlation_id=payload.correlation_id,
            status=JobStatus.SUCCEEDED,
            stub=stub,
            message=f"Scenario graded via {provider}",
            total_score=float(data.get("total_score") or 0),
            max_total_score=float(data.get("max_total_score") or 100),
            passed=data.get("passed"),
            feedback_markdown=data.get("feedback_markdown"),
        )


class InterviewLockValidateService:
    async def validate(
        self, payload: InterviewValidateLockRequest
    ) -> InterviewValidateLockResponse:
        issues: list[str] = []
        warnings: list[str] = []
        if len(payload.questions) < 3:
            issues.append("Need at least 3 questions before lock")
        empty = [
            str(q.get("id") or i)
            for i, q in enumerate(payload.questions)
            if isinstance(q, dict) and not (q.get("stem") or q.get("text"))
        ]
        if empty:
            issues.append(f"Empty question stems: {', '.join(empty[:5])}")
        for mid in payload.mandatory_question_ids:
            ids = {
                str(q.get("id"))
                for q in payload.questions
                if isinstance(q, dict) and q.get("id")
            }
            if mid not in ids:
                issues.append(f"Mandatory question missing: {mid}")
        if len(payload.questions) > 25:
            warnings.append("Large question set may increase interview fatigue")
        return InterviewValidateLockResponse(
            correlation_id=payload.correlation_id,
            status=JobStatus.SUCCEEDED,
            stub=True,
            message="Interview lock validation complete",
            ready_to_lock=len(issues) == 0,
            issues=issues,
            warnings=warnings,
        )


class FollowUpLogicService:
    async def define(self, payload: FollowUpLogicRequest) -> FollowUpLogicResponse:
        system = (
            "Define adaptive follow-up logic for an interview pack. Return JSON: "
            "{\"rules\":[{\"question_id\",\"trigger\",\"follow_up_prompt\",\"max_depth\"}],"
            "\"logic_markdown\"}"
        )
        user = (
            f"Role: {payload.role_title}\nPersona: {payload.persona.value}\n"
            f"Questions: {payload.questions}"
        )

        def local() -> dict:
            rules = []
            for i, q in enumerate(payload.questions):
                if not isinstance(q, dict):
                    continue
                qid = str(q.get("id") or f"q-{i + 1}")
                stem = str(q.get("stem") or q.get("text") or "this topic")
                rules.append(
                    {
                        "question_id": qid,
                        "trigger": "score_below_0.6_or_vague_answer",
                        "follow_up_prompt": (
                            f"Probe deeper on: {stem[:120]}. Ask for a concrete example, "
                            "metrics, and trade-offs."
                        ),
                        "max_depth": 2,
                    }
                )
            return {
                "rules": rules,
                "logic_markdown": (
                    "# Follow-up logic\n\n"
                    "- If answer score < 0.6 → ask clarifying example follow-up\n"
                    "- If answer is strong → optional stretch follow-up (depth 1)\n"
                ),
            }

        data, stub, provider = await generate_json(
            system=system, user=user, local_factory=local
        )
        rules = [
            FollowUpRule(
                question_id=str(r.get("question_id")),
                trigger=str(r.get("trigger") or "score_below_0.6"),
                follow_up_prompt=str(r.get("follow_up_prompt") or ""),
                max_depth=int(r.get("max_depth") or 1),
            )
            for r in data.get("rules") or []
            if isinstance(r, dict) and r.get("question_id") and r.get("follow_up_prompt")
        ]
        return FollowUpLogicResponse(
            correlation_id=payload.correlation_id,
            status=JobStatus.SUCCEEDED,
            stub=stub,
            message=f"Follow-up logic via {provider}",
            rules=rules,
            logic_markdown=data.get("logic_markdown"),
        )


class MergeCustomQuestionsService:
    async def merge(
        self, payload: MergeCustomQuestionsRequest
    ) -> MergeCustomQuestionsResponse:
        merged: list[dict] = []
        seen: set[str] = set()
        for q in list(payload.existing_questions) + list(payload.custom_questions):
            if not isinstance(q, dict):
                continue
            qid = str(q.get("id") or f"cq-{len(merged) + 1}")
            if qid in seen:
                continue
            seen.add(qid)
            row = dict(q)
            row["id"] = qid
            if q in payload.custom_questions:
                meta = row.get("metadata")
                if not isinstance(meta, dict):
                    meta = {}
                meta["source"] = "company_custom"
                row["metadata"] = meta
            merged.append(row)
        mandatory = list(payload.mandatory_question_ids)
        for q in payload.custom_questions:
            if isinstance(q, dict) and q.get("mandatory") and q.get("id"):
                mid = str(q["id"])
                if mid not in mandatory:
                    mandatory.append(mid)
        return MergeCustomQuestionsResponse(
            correlation_id=payload.correlation_id,
            status=JobStatus.SUCCEEDED,
            stub=True,
            message="Custom company questions merged into interview pack",
            questions=merged,
            mandatory_question_ids=mandatory,
            added_count=len(payload.custom_questions),
        )
