File size: 3,792 Bytes
eeff6a0
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
import json
import logging
import os

import httpx
from google import genai
from google.genai import errors, types
from pydantic import BaseModel, Field
from sqlalchemy.orm import Session

import app.config  # Loads Backend/.env
from app.models.department import Department

logger = logging.getLogger(__name__)


class ClassificationResult(BaseModel):
    department_id: int | None
    confidence: float = Field(ge=0, le=1)


def classify_query(
    db: Session,
    message: str,
) -> tuple[Department, float] | None:
    """Return a validated Gemini result, or None to use keyword routing."""
    gemini_keys = [
    os.getenv("GEMINI_API_KEY_1"),
    os.getenv("GEMINI_API_KEY_2"),
    os.getenv("GEMINI_API_KEY_3"),
    os.getenv("GEMINI_API_KEY_4"),
    os.getenv("GEMINI_API_KEY_5"),
    os.getenv("GEMINI_API_KEY_6"),
    os.getenv("GEMINI_API_KEY_7"),
    os.getenv("GEMINI_API_KEY_8"),
    os.getenv("GEMINI_API_KEY_9"),
]


    gemini_keys = [
    key for key in gemini_keys if key
]


    api_key = gemini_keys[0] if gemini_keys else None
    if not api_key:
        logger.warning("Gemini key missing; using keyword routing.")
        return None

    departments = (
        db.query(Department)
        .filter(Department.is_active.is_(True))
        .all()
    )
    if not departments:
        return None

    department_data = [
        {
            "id": department.id,
            "name": department.name,
            "code": department.code,
            "description": department.description or "",
            "keywords": department.keywords or "",
            "example_queries": department.query or "",
        }
        for department in departments
    ]

    try:
        with genai.Client(
            api_key=api_key,
            http_options=types.HttpOptions(
                timeout=15000,
                retry_options=types.HttpRetryOptions(attempts=1),
            ),
        ) as client:
            response = client.models.generate_content(
                  model="gemini-3.6-flash",
                contents=json.dumps({
                    "departments": department_data,
                    "student_message": message,
                }),
                config=types.GenerateContentConfig(
                    system_instruction=(
                     "You are an AI query routing assistant for  a university. "
                      "Classify student queries and select the most suitable department. "
                      "Use only the supplied department information. "
                      "Return only the department_id and confidence score. "
                      "Do not answer the student query."
),
                    response_mime_type="application/json",
                    response_schema=ClassificationResult,
                    automatic_function_calling=(
                        types.AutomaticFunctionCallingConfig(disable=True)
                    ),
                ),
            )

        result = ClassificationResult.model_validate_json(
            response.text or ""
        )

        department = next(
            (
                department
                for department in departments
                if department.id == result.department_id
            ),
            None,
        )

        if department is None:
            logger.warning(
                "Gemini returned no valid department; using keyword routing."
            )
            return None

        return department, result.confidence

    except (errors.APIError, httpx.HTTPError, ValueError) as exc:
        # Log the error type without exposing credentials or query content.
        logger.warning(
            "Gemini classification failed (%s); using keyword routing.",
            type(exc).__name__,
        )
        return None