#!/usr/bin/env python3
"""
Log Receiver Daemon — TCP syslog 接收器 (port 5140)

客戶端 rsyslog 設定：
  /etc/rsyslog.d/99-portal.conf:
    *.* action(type="omfwd" target="portal.steps.tw" port="5140" protocol="tcp"
               queue.type="linkedList" queue.size="10000" resumeRetryCount="-1")

身份識別：以來源 IP 對應 portal 資料庫的 host.ip_address。
緩衝策略：每個來源 IP 獨立緩衝，每 FLUSH_INTERVAL 秒或累積 MAX_LINES 行時 POST 到 portal。
"""

import asyncio
import http.client
import json
import logging
import os
import signal
import sys
from collections import defaultdict

# ── 設定 ──────────────────────────────────────────────────────────────────────
LISTEN_HOST    = "0.0.0.0"
LISTEN_PORT    = 5140
PORTAL_URL     = "http://127.0.0.1/api/logs/internal"
PORTAL_HOST    = "portal.steps.tw"   # Apache vhost Host header
INTERNAL_KEY   = ""                   # 從 .env 讀入
FLUSH_INTERVAL = 300                  # 預設 5 分鐘，可用 .env LOG_FLUSH_INTERVAL 覆蓋
MAX_LINES      = 2000                 # 累積到這麼多行就立即 flush
MAX_BYTES      = 1024 * 1024          # 1MB 上限，超過就 flush

LOG_FILE = "/var/www/html/portal/storage/logs/log_receiver.log"

logging.basicConfig(
    level=logging.INFO,
    format="%(asctime)s [%(levelname)s] %(message)s",
    handlers=[
        logging.FileHandler(LOG_FILE),
        logging.StreamHandler(sys.stdout),
    ],
)
log = logging.getLogger("log_receiver")

# ── 讀取 .env 的 LOG_RECEIVER_KEY ─────────────────────────────────────────────
def load_env_key() -> tuple[str, int]:
    env_path = "/var/www/html/portal/.env"
    key      = ""
    interval = 300
    try:
        with open(env_path) as f:
            for line in f:
                line = line.strip()
                if line.startswith("LOG_RECEIVER_KEY="):
                    key = line.split("=", 1)[1].strip().strip('"').strip("'")
                elif line.startswith("LOG_FLUSH_INTERVAL="):
                    try:
                        interval = max(30, int(line.split("=", 1)[1].strip()))
                    except ValueError:
                        pass
    except Exception as e:
        log.warning(f"Cannot read .env: {e}")
    return key, interval

# ── 每個 IP 的緩衝區 ─────────────────────────────────────────────────────────
# buffer key = (source_ip, log_type)
buffers:      dict[tuple, list[str]] = defaultdict(list)
buffer_bytes: dict[tuple, int]       = defaultdict(int)

TAG_TYPE_MAP = {
    "apache-access":  "apache_access",
    "apache_access":  "apache_access",
    "apache-error":   "apache_error",
    "apache_error":   "apache_error",
    "httpd-access":   "apache_access",
    "httpd-error":    "apache_error",
    "nginx-access":   "nginx_access",
    "nginx_access":   "nginx_access",
    "nginx-error":    "nginx_error",
    "nginx_error":    "nginx_error",
    "auth":           "auth",
    "sshd":           "auth",
    "sudo":           "auth",
    "dmesg":          "kernel",
    "kernel":         "kernel",
}

def detect_log_type(line: str) -> str:
    """從 syslog TAG 欄位判斷 log 類型"""
    import re
    # 標準格式：<pri>timestamp hostname TAG[pid]: msg  OR  imfile 格式：TAG msg（無冒號）
    m = re.match(
        r'^(?:<\d+>)?(?:\w{3}\s+\d+\s+\d{2}:\d{2}:\d{2}\s+)?\S+\s+(\S+?)(?:\[\d+\])?(?::|(?=\s))',
        line
    )
    if m:
        tag = m.group(1).lower().rstrip(':').strip()
        if tag in TAG_TYPE_MAP:
            return TAG_TYPE_MAP[tag]
        for key, val in TAG_TYPE_MAP.items():
            if key in tag:
                return val
    return "syslog"

def flush_buffer(source_ip: str, log_type: str | None = None):
    """flush 指定 (ip, type)，若 log_type 為 None 則 flush 該 ip 所有 type"""
    keys = [(source_ip, log_type)] if log_type else [k for k in buffers if k[0] == source_ip]
    for key in keys:
        lines = buffers.get(key)
        if not lines:
            continue
        content = "\n".join(lines)
        buffers[key] = []
        buffer_bytes[key] = 0
        post_to_portal(source_ip, key[1], content)

def post_to_portal(source_ip: str, log_type: str, content: str):
    payload = json.dumps({
        "source_ip": source_ip,
        "log_type":  log_type,
        "lines":     content,
    }).encode("utf-8")

    try:
        conn = http.client.HTTPConnection("127.0.0.1", 80, timeout=10)
        conn.request("POST", "/api/logs/internal", body=payload, headers={
            "Content-Type":   "application/json",
            "Host":           PORTAL_HOST,
            "X-Internal-Key": INTERNAL_KEY,
        })
        r    = conn.getresponse()
        body = r.read().decode("utf-8", errors="replace")
        conn.close()
        if r.status == 200:
            resp = json.loads(body)
            log.info(f"Flushed {source_ip}: collection_id={resp.get('id')} ({len(content)} bytes)")
        else:
            log.error(f"Portal HTTP {r.status} for {source_ip}: {body[:200]}")
    except Exception:
        log.exception(f"Portal post failed for {source_ip}")

# ── 週期性 flush（每 FLUSH_INTERVAL 秒） ──────────────────────────────────────
async def periodic_flush():
    while True:
        await asyncio.sleep(FLUSH_INTERVAL)
        done_ips = {k[0] for k in list(buffers.keys())}
        for ip in done_ips:
            flush_buffer(ip)

# ── 每個連線的處理 ────────────────────────────────────────────────────────────
async def handle_client(reader: asyncio.StreamReader, writer: asyncio.StreamWriter):
    addr      = writer.get_extra_info("peername")
    source_ip = addr[0] if addr else "0.0.0.0"
    log.info(f"Connected: {source_ip}")

    try:
        while True:
            line = await asyncio.wait_for(reader.readline(), timeout=300)
            if not line:
                break
            decoded = line.decode("utf-8", errors="replace").rstrip("\r\n")
            if not decoded:
                continue

            log_type = detect_log_type(decoded)
            key      = (source_ip, log_type)
            buffers[key].append(decoded)
            buffer_bytes[key] += len(decoded)

            # 達到閾值立即 flush 該 type
            if len(buffers[key]) >= MAX_LINES or buffer_bytes[key] >= MAX_BYTES:
                flush_buffer(source_ip, log_type)

    except asyncio.TimeoutError:
        log.info(f"Timeout (300s idle): {source_ip}")
    except asyncio.IncompleteReadError:
        pass
    except Exception as e:
        log.warning(f"Client {source_ip} error: {e}")
    finally:
        # 連線結束時 flush 該 ip 所有 type
        flush_buffer(source_ip)
        try:
            writer.close()
        except Exception:
            pass
        log.info(f"Disconnected: {source_ip}")

# ── 主程式 ────────────────────────────────────────────────────────────────────
async def main():
    global INTERNAL_KEY, FLUSH_INTERVAL
    INTERNAL_KEY, FLUSH_INTERVAL = load_env_key()
    if not INTERNAL_KEY:
        log.error("LOG_RECEIVER_KEY not set in .env — daemon will fail auth checks")
    else:
        log.info(f"Internal key loaded ({len(INTERNAL_KEY)} chars), flush interval={FLUSH_INTERVAL}s")

    server = await asyncio.start_server(handle_client, LISTEN_HOST, LISTEN_PORT)
    log.info(f"Listening on {LISTEN_HOST}:{LISTEN_PORT}")

    flush_task = asyncio.create_task(periodic_flush())
    loop = asyncio.get_running_loop()

    def _shutdown():
        log.info("Shutting down — flushing all buffers…")
        for ip in list(buffers.keys()):
            flush_buffer(ip)
        flush_task.cancel()
        server.close()

    for sig in (signal.SIGINT, signal.SIGTERM):
        loop.add_signal_handler(sig, _shutdown)

    async with server:
        await server.serve_forever()

if __name__ == "__main__":
    asyncio.run(main())
