File size: 6,547 Bytes
c4f5819
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
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
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
"""Retrieval-augmented Q&A over the stored log events.

Offline path (always available): keyword/TF-IDF ranking over rows fetched from
the DB + a rule-based Thai summary. Optional path: when OPENROUTER_API_KEY is
set, the retrieved context is sent to an LLM for a richer answer. Any LLM error
falls back to the local summary so the system never hard-fails.
"""
from __future__ import annotations

import math
from collections import Counter, defaultdict
from typing import Any

import httpx

from .classifier import CATEGORY_META, detect_rule, expand_query, row_text, tokenize
from .config import get_settings
from .db import search_rows


async def retrieve(query: str, limit: int = 8) -> list[dict[str, Any]]:
    """Keyword retrieval: pull candidate rows then TF-IDF rank in-process."""
    terms, category_terms = expand_query(query)
    # Pull a candidate pool (filtered by the raw query when possible).
    pool = await search_rows(query, limit=600)
    if not pool:
        pool = await search_rows("", limit=600)
    if not terms or not pool:
        return [_as_hit(r, 1.0) for r in pool[:limit]]

    # Build a tiny TF-IDF index over the candidate pool.
    docs = [Counter(tokenize(row_text(r))) for r in pool]
    doc_freq: Counter[str] = Counter()
    for d in docs:
        doc_freq.update(d.keys())
    total = max(len(docs), 1)
    idf = {t: math.log((1 + total) / (1 + df)) + 1 for t, df in doc_freq.items()}

    query_counts = Counter(terms)
    scores: defaultdict[int, float] = defaultdict(float)
    for idx, doc in enumerate(docs):
        for term, qtf in query_counts.items():
            tf = doc.get(term, 0)
            if tf:
                scores[idx] += (1 + math.log(tf)) * idf.get(term, 1.0) * qtf
        if pool[idx].get("category") in category_terms:
            scores[idx] += 12
        if query.lower() in row_text(pool[idx]).lower():
            scores[idx] += 8

    ranked = sorted(scores.items(), key=lambda kv: kv[1], reverse=True)[:limit]
    return [_as_hit(pool[idx], round(score, 3)) for idx, score in ranked]


def _as_hit(row: dict[str, Any], score: float) -> dict[str, Any]:
    return {"score": score, "row": row, "rule": detect_rule(row), "retrieval": "keyword"}


def answer_local(question: str, hits: list[dict[str, Any]]) -> dict[str, Any]:
    rows = [h["row"] for h in hits]
    by_cat = Counter(r.get("category") for r in rows)
    by_sev = Counter(r.get("severity") for r in rows)
    if not rows:
        answer = ("## ยังไม่พบหลักฐาน\n\n"
                  "- ลองถามด้วย IP, category เช่น `brute_force`, `c2_beacon`\n"
                  "- หรืออัปโหลด/ส่ง mock log ก่อนแล้วถามใหม่")
    else:
        top = rows[0]
        meta = CATEGORY_META.get(top.get("category"), CATEGORY_META["benign"])
        cats = ", ".join(c for c, _ in by_cat.most_common(3) if c)
        answer = (
            f"## สรุปความเสี่ยง\n\n"
            f"- พบหลักฐานที่เกี่ยวข้อง **{len(rows)} รายการ**\n"
            f"- กลุ่มเด่น: **{cats}**\n"
            f"- รายการสำคัญสุด: **{meta['label']}** severity=`{top.get('severity')}` source=`{top.get('source')}`\n"
            f"- เวลา: `{top.get('ts')}`\n\n"
            f"## เหตุผล\n\n- {', '.join(detect_rule(top)['reasons'])}\n\n"
            f"## Next action\n\n- {meta['runbook']}"
        )
    return {"answer": answer, "summary": {"categories": dict(by_cat), "severities": dict(by_sev)}}


def build_context(hits: list[dict[str, Any]]) -> str:
    lines = []
    for i, hit in enumerate(hits, 1):
        row = hit["row"]
        meta = CATEGORY_META.get(row.get("category"), CATEGORY_META["benign"])
        lines.append("\n".join([
            f"[{i}] score={hit['score']}",
            f"time={row.get('ts')} source={row.get('source')} severity={row.get('severity')}",
            f"category={row.get('category')} label={meta['label']} mitre={meta['mitre']}",
            f"src_ip={row.get('src_ip', '-')} dst_ip={row.get('dst_ip', '-')}",
            f"message={row.get('message')}",
            f"rule_reason={'; '.join(hit['rule']['reasons'])}",
            f"runbook={meta['runbook']}",
        ]))
    return "\n\n".join(lines)


async def call_llm(question: str, context: str, model: str = "") -> str:
    settings = get_settings()
    if not settings.openrouter_api_key:
        raise RuntimeError("OPENROUTER_API_KEY is not set")
    payload = {
        "model": model or settings.openrouter_model,
        "messages": [
            {"role": "system", "content": (
                "You are a defensive security log analyst. Answer in Thai. "
                "Use only the provided retrieved logs and runbook context. "
                "Format as concise Markdown with sections: สรุป, หลักฐาน, ความเสี่ยง, Next action. "
                "Use short bullet points and inline code for IPs, users, MITRE IDs, and log fields. "
                "Do not provide offensive instructions.")},
            {"role": "user", "content": f"คำถาม: {question}\n\nRetrieved context:\n{context}"},
        ],
        "temperature": 0.2,
        "max_tokens": 700,
    }
    async with httpx.AsyncClient(timeout=45) as client:
        resp = await client.post(
            "https://openrouter.ai/api/v1/chat/completions",
            headers={"Authorization": f"Bearer {settings.openrouter_api_key}",
                     "Content-Type": "application/json"},
            json=payload,
        )
        resp.raise_for_status()
        data = resp.json()
    return data["choices"][0]["message"]["content"]


async def ask(question: str, limit: int = 8, use_llm: bool = False, model: str = "") -> dict[str, Any]:
    hits = await retrieve(question, limit)
    local = answer_local(question, hits)
    answer, mode, llm_error = local["answer"], "local", None
    if use_llm and get_settings().llm_ready:
        try:
            answer = await call_llm(question, build_context(hits), model)
            mode = "openrouter"
        except Exception as exc:  # keep the demo alive on rate-limit/errors
            llm_error = str(exc)
    return {
        "answer": answer, "mode": mode, "llm_error": llm_error,
        "retrieval": hits[0]["retrieval"] if hits else "none",
        "local_summary": local["summary"],
        "hits": hits,
    }