# -*- coding: utf-8 -*-
"""
ATLUS 搭載PC に常駐させる取得エージェント。

役割は「ATLUS から必要なものを受け取って渡す」ところまで。
案件の積算そのものを ATLUS に完走させる作りにはしない。計算はクラウド側で行う。
ここを守らないと、処理量が ATLUS 搭載PC 1台の速さで頭打ちになり、
その1台が止まった日に入札が止まる。

かならず「ログオン中のユーザーのセッション」で起動すること。
SSH 越しに起動すると対話デスクトップに触れないため、画面は真っ黒になり
ウィンドウ操作も一切効かない。

  下見のとき   py -3 atlas_agent.py --token <合言葉> --allow 192.168.1.50
  普段まわすとき py -3 atlas_agent.py --token <合言葉> --allow 192.168.1.50 \
                   --no-exec --root C:\\atlas\\out

必要なもの:
  pip install pillow          画面の取得に要る
  pip install pywinauto       ATLUS のコントロールを触るのに要る（無くても起動はする）
"""

import argparse
import ctypes
import glob as globmod
import hashlib
import io
import json
import os
import socket
import subprocess
import sys
import time
import traceback
from ctypes import wintypes
from contextlib import redirect_stdout
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
from urllib.parse import urlparse, parse_qs, unquote

VERSION = "1.0"

# 画面の座標とスクリーンショットをずらさないため、拡大率の扱いを先に固定する
try:
    ctypes.windll.shcore.SetProcessDpiAwareness(2)
except Exception:
    try:
        ctypes.windll.user32.SetProcessDPIAware()
    except Exception:
        pass

user32 = ctypes.windll.user32
kernel32 = ctypes.windll.kernel32

CFG = {"token": "", "allow": [], "root": "", "log": None, "no_exec": False}


def log(msg):
    line = "%s  %s" % (time.strftime("%Y-%m-%d %H:%M:%S"), msg)
    print(line, flush=True)
    if CFG["log"]:
        try:
            with open(CFG["log"], "a", encoding="utf-8") as f:
                f.write(line + "\n")
        except Exception:
            pass


# ---------------------------------------------------------------- ウィンドウ一覧

def _s(text):
    """題名に壊れた文字（対になっていないサロゲート）が混じることがあるので落とす。
    そのまま JSON に載せると受け取る側が読めない。"""
    return text.encode("utf-8", "replace").decode("utf-8")


def list_windows(visible_only=True):
    """pywinauto が無くても動く、標準機能だけのウィンドウ一覧。"""
    out = []
    WNDENUMPROC = ctypes.WINFUNCTYPE(wintypes.BOOL, wintypes.HWND, wintypes.LPARAM)

    def cb(hwnd, _lparam):
        if visible_only and not user32.IsWindowVisible(hwnd):
            return True
        n = user32.GetWindowTextLengthW(hwnd)
        buf = ctypes.create_unicode_buffer(n + 1)
        user32.GetWindowTextW(hwnd, buf, n + 1)
        title = _s(buf.value)
        if visible_only and not title:
            return True
        cls = ctypes.create_unicode_buffer(256)
        user32.GetClassNameW(hwnd, cls, 256)
        pid = wintypes.DWORD()
        user32.GetWindowThreadProcessId(hwnd, ctypes.byref(pid))
        r = wintypes.RECT()
        user32.GetWindowRect(hwnd, ctypes.byref(r))
        out.append({
            "hwnd": int(hwnd),
            "title": title,
            "class": _s(cls.value),
            "pid": int(pid.value),
            "rect": [r.left, r.top, r.right, r.bottom],
            "visible": bool(user32.IsWindowVisible(hwnd)),
        })
        return True

    user32.EnumWindows(WNDENUMPROC(cb), 0)
    return out


def find_window(pattern):
    """題名の部分一致でウィンドウを1つ選ぶ。見つからなければ None。"""
    if not pattern:
        return None
    pat = pattern.lower()
    hits = [w for w in list_windows() if pat in w["title"].lower()]
    if not hits:
        return None
    # 大きいものを本体とみなす（小さい浮きものを拾わないため）
    hits.sort(key=lambda w: (w["rect"][2] - w["rect"][0]) * (w["rect"][3] - w["rect"][1]),
              reverse=True)
    return hits[0]


# ---------------------------------------------------------------- 画面の取得

def grab_png(bbox=None, scale=1.0):
    from PIL import ImageGrab
    img = ImageGrab.grab(bbox=bbox, all_screens=True)
    if scale and scale != 1.0:
        w = max(1, int(img.width * scale))
        h = max(1, int(img.height * scale))
        img = img.resize((w, h))
    buf = io.BytesIO()
    img.save(buf, format="PNG", optimize=True)
    return buf.getvalue(), img.size


# ---------------------------------------------------------------- pywinauto

_uia = {"desktop": None}


def uia_desktop():
    from pywinauto import Desktop
    if _uia["desktop"] is None:
        _uia["desktop"] = Desktop(backend="uia")
    return _uia["desktop"]


def uia_window(hwnd=None, title=None):
    d = uia_desktop()
    if hwnd:
        return d.window(handle=int(hwnd))
    w = find_window(title)
    if not w:
        raise ValueError("そのウィンドウが見つからない: %r" % title)
    return d.window(handle=w["hwnd"])


def uia_child(win, spec):
    """spec は auto_id / title / control_type / index を任意に組み合わせた辞書。"""
    crit = {}
    if spec.get("auto_id"):
        crit["auto_id"] = spec["auto_id"]
    if spec.get("title"):
        crit["title"] = spec["title"]
    if spec.get("title_re"):
        crit["title_re"] = spec["title_re"]
    if spec.get("control_type"):
        crit["control_type"] = spec["control_type"]
    if spec.get("class_name"):
        crit["class_name"] = spec["class_name"]
    if not crit:
        return win
    ctrl = win.child_window(**crit)
    if spec.get("index"):
        ctrl = ctrl[int(spec["index"])]
    return ctrl


# ---------------------------------------------------------------- HTTP

class Handler(BaseHTTPRequestHandler):
    server_version = "AtlasAgent/" + VERSION
    protocol_version = "HTTP/1.1"

    def log_message(self, fmt, *args):
        pass  # 既定の標準エラー出力を止め、log() に一本化する

    # -- 送り返し ------------------------------------------------
    def _send(self, code, body, ctype="application/json; charset=utf-8"):
        if isinstance(body, (dict, list)):
            body = json.dumps(body, ensure_ascii=False, indent=1).encode("utf-8")
        elif isinstance(body, str):
            body = body.encode("utf-8")
        self.send_response(code)
        self.send_header("Content-Type", ctype)
        self.send_header("Content-Length", str(len(body)))
        self.end_headers()
        self.wfile.write(body)

    def _err(self, code, msg):
        log("  -> %d %s" % (code, msg))
        self._send(code, {"ok": False, "error": msg})

    # -- 入口の見張り --------------------------------------------
    def _guard(self):
        ip = self.client_address[0]
        if CFG["allow"] and ip not in CFG["allow"]:
            self._err(403, "この接続元は許していない: %s" % ip)
            return False
        if self.headers.get("X-Token", "") != CFG["token"]:
            self._err(401, "合言葉が違う")
            return False
        return True

    def _body(self):
        n = int(self.headers.get("Content-Length") or 0)
        return self.rfile.read(n) if n else b""

    def _json(self):
        b = self._body()
        return json.loads(b.decode("utf-8")) if b else {}

    def _safe_path(self, p):
        """--root を指定したときは、その下から出さない。"""
        p = os.path.abspath(p)
        root = CFG["root"]
        if root and not p.lower().startswith(os.path.abspath(root).lower()):
            raise ValueError("--root の外は触れない: %s" % p)
        return p

    # -- 振り分け ------------------------------------------------
    def do_GET(self):
        self._route("GET")

    def do_POST(self):
        self._route("POST")

    def _route(self, method):
        u = urlparse(self.path)
        q = {k: v[0] for k, v in parse_qs(u.query).items()}
        path = u.path.rstrip("/") or "/"
        if not self._guard():
            return
        log("%s %s %s" % (self.client_address[0], method, self.path[:200]))
        try:
            fn = getattr(self, "h_" + path.strip("/").replace("/", "_") or "h_", None)
            if fn is None:
                self._err(404, "そんな口は無い: %s" % path)
                return
            fn(q)
        except Exception as e:
            log(traceback.format_exc())
            self._err(500, "%s: %s" % (type(e).__name__, e))

    # -- それぞれの口 --------------------------------------------
    def h_(self, q):
        self.h_health(q)

    def h_health(self, q):
        caps = {}
        for m in ("PIL", "pywinauto"):
            try:
                __import__(m)
                caps[m] = True
            except Exception:
                caps[m] = False
        sid = wintypes.DWORD()
        kernel32.ProcessIdToSessionId(kernel32.GetCurrentProcessId(), ctypes.byref(sid))
        self._send(200, {
            "ok": True,
            "version": VERSION,
            "host": socket.gethostname(),
            "user": os.environ.get("USERNAME"),
            "session_id": int(sid.value),
            "interactive": int(sid.value) != 0,
            "screen": [user32.GetSystemMetrics(0), user32.GetSystemMetrics(1)],
            "capabilities": caps,
            "root": CFG["root"] or None,
            "cwd": os.getcwd(),
        })

    def h_windows(self, q):
        vis = q.get("all") != "1"
        ws = list_windows(visible_only=vis)
        pat = (q.get("q") or "").lower()
        if pat:
            ws = [w for w in ws if pat in w["title"].lower()]
        self._send(200, {"ok": True, "count": len(ws), "windows": ws})

    def h_shot(self, q):
        bbox = None
        if q.get("window"):
            w = find_window(q["window"])
            if not w:
                self._err(404, "そのウィンドウが見つからない: %s" % q["window"])
                return
            bbox = tuple(w["rect"])
        elif q.get("bbox"):
            bbox = tuple(int(x) for x in q["bbox"].split(","))
        png, size = grab_png(bbox, float(q.get("scale") or 1.0))
        self.send_response(200)
        self.send_header("Content-Type", "image/png")
        self.send_header("Content-Length", str(len(png)))
        self.send_header("X-Image-Size", "%dx%d" % size)
        self.end_headers()
        self.wfile.write(png)

    def h_exec(self, q):
        if CFG["no_exec"]:
            self._err(403, "--no-exec で起動しているので、任意コマンドの実行は閉じている")
            return
        d = self._json()
        argv = d.get("argv")
        cmd = d.get("cmd")
        if not argv and not cmd:
            self._err(400, "argv か cmd のどちらかが要る")
            return
        t0 = time.time()
        try:
            p = subprocess.run(
                argv if argv else cmd,
                shell=bool(cmd),
                cwd=d.get("cwd") or None,
                capture_output=True,
                timeout=float(d.get("timeout") or 120),
            )
            out = p.stdout.decode("cp932", "replace")
            err = p.stderr.decode("cp932", "replace")
            rc = p.returncode
            timed_out = False
        except subprocess.TimeoutExpired as e:
            out = (e.stdout or b"").decode("cp932", "replace")
            err = (e.stderr or b"").decode("cp932", "replace")
            rc = None
            timed_out = True
        self._send(200, {"ok": rc == 0, "rc": rc, "timed_out": timed_out,
                         "elapsed": round(time.time() - t0, 2),
                         "stdout": out, "stderr": err})

    def h_start(self, q):
        """待たずに起動だけする（ATLUS 本体の立ち上げ用）。"""
        if CFG["no_exec"]:
            self._err(403, "--no-exec で起動しているので、任意コマンドの実行は閉じている")
            return
        d = self._json()
        argv = d.get("argv") or d.get("cmd")
        p = subprocess.Popen(argv, shell=isinstance(argv, str), cwd=d.get("cwd") or None)
        self._send(200, {"ok": True, "pid": p.pid})

    def h_upload(self, q):
        path = self._safe_path(unquote(q.get("path") or ""))
        os.makedirs(os.path.dirname(path) or ".", exist_ok=True)
        data = self._body()
        with open(path, "wb") as f:
            f.write(data)
        self._send(200, {"ok": True, "path": path, "bytes": len(data)})

    def h_download(self, q):
        path = self._safe_path(unquote(q.get("path") or ""))
        with open(path, "rb") as f:
            data = f.read()
        # 何を、いつ作られたものを持ち出したかを、受け取る側にも記録にも残す。
        # クラウド側はこの2つを行に刻み、どの版から出た数字かを後から辿れるようにする。
        sha = hashlib.sha256(data).hexdigest()
        mtime = time.strftime("%Y-%m-%dT%H:%M:%S", time.localtime(os.path.getmtime(path)))
        log("  持ち出し %s  %dB  sha=%s  作成=%s" % (path, len(data), sha[:16], mtime))
        self.send_response(200)
        self.send_header("Content-Type", "application/octet-stream")
        self.send_header("Content-Length", str(len(data)))
        self.send_header("X-Sha256", sha)
        self.send_header("X-Mtime", mtime)
        self.send_header("X-Source-Host", socket.gethostname())
        self.end_headers()
        self.wfile.write(data)

    def h_wait(self, q):
        """書き出しが終わるまで待つ。

        ATLUS の Excel 出力は、ファイルが現れてから中身が揃うまでに間がある。
        現れた瞬間に掴むと、途中まで書かれたものを持ち出して静かに壊れる。
        大きさが動かなくなってから返す。
        """
        pattern = unquote(q.get("glob") or "")
        path = unquote(q.get("path") or "")
        timeout = float(q.get("timeout") or 300)
        stable = float(q.get("stable") or 2.0)
        t0 = time.time()
        last, since, found = -1, None, None
        while time.time() - t0 < timeout:
            if pattern:
                hits = [p for p in globmod.glob(pattern) if os.path.isfile(p)]
                found = max(hits, key=os.path.getmtime) if hits else None
            else:
                found = path if path and os.path.exists(path) else None
            if found:
                size = os.path.getsize(found)
                if size != last:
                    last, since = size, None      # まだ増えている
                elif size > 0:
                    if since is None:
                        since = time.time()       # 止まった。ここから数える
                    elif time.time() - since >= stable:
                        self._send(200, {"ok": True, "path": self._safe_path(found),
                                         "bytes": size,
                                         "waited": round(time.time() - t0, 1)})
                        return
            time.sleep(0.4)
        self._send(200, {"ok": False, "error": "時間内に書き上がらなかった",
                         "path": found, "bytes": last if last >= 0 else None,
                         "waited": round(time.time() - t0, 1)})

    def h_ls(self, q):
        path = self._safe_path(unquote(q.get("path") or "."))
        rows = []
        for name in sorted(os.listdir(path)):
            fp = os.path.join(path, name)
            try:
                st = os.stat(fp)
                rows.append({"name": name, "dir": os.path.isdir(fp), "bytes": st.st_size,
                             "mtime": time.strftime("%Y-%m-%d %H:%M", time.localtime(st.st_mtime))})
            except OSError:
                rows.append({"name": name, "error": True})
        self._send(200, {"ok": True, "path": path, "entries": rows})

    def h_tree(self, q):
        """ATLUS のコントロールツリーを出す。自動化できるかはここで決まる。"""
        win = uia_window(q.get("hwnd"), q.get("window"))
        buf = io.StringIO()
        with redirect_stdout(buf):
            win.print_control_identifiers(depth=int(q.get("depth") or 3))
        self._send(200, buf.getvalue(), "text/plain; charset=utf-8")

    def h_ui(self, q):
        """focus / click / type / key / set_text / get_text"""
        d = self._json()
        action = d.get("action")
        win = uia_window(d.get("hwnd"), d.get("window"))
        win.set_focus()
        if action == "focus":
            res = "focused"
        elif action == "key":
            from pywinauto.keyboard import send_keys
            send_keys(d["keys"])
            res = "keys sent"
        else:
            ctrl = uia_child(win, d.get("target") or {})
            if action == "click":
                ctrl.click_input(button=d.get("button", "left"),
                                 double=bool(d.get("double")))
                res = "clicked"
            elif action == "type":
                ctrl.type_keys(d["text"], with_spaces=True, with_newlines=True)
                res = "typed"
            elif action == "set_text":
                ctrl.set_edit_text(d["text"])
                res = "set"
            elif action == "get_text":
                res = ctrl.window_text()
            else:
                self._err(400, "知らない action: %r" % action)
                return
        self._send(200, {"ok": True, "action": action, "result": res})


def main():
    ap = argparse.ArgumentParser()
    ap.add_argument("--token", required=True, help="合言葉。X-Token ヘッダで照合する")
    ap.add_argument("--host", default="0.0.0.0")
    ap.add_argument("--port", type=int, default=8765)
    ap.add_argument("--allow", action="append", default=[],
                    help="つなげてよい相手の IP。何度でも書ける。既定は全許可なので必ず指定する")
    ap.add_argument("--root", default="", help="ファイル操作をこの下に閉じ込める")
    ap.add_argument("--log", default="", help="記録の書き出し先")
    ap.add_argument("--no-exec", action="store_true",
                    help="任意コマンドの実行を閉じる。下見が済んだら付けて回す")
    a = ap.parse_args()

    CFG["token"] = a.token
    CFG["allow"] = a.allow
    CFG["root"] = a.root
    CFG["log"] = a.log or None
    CFG["no_exec"] = a.no_exec

    sid = wintypes.DWORD()
    kernel32.ProcessIdToSessionId(kernel32.GetCurrentProcessId(), ctypes.byref(sid))
    log("ATLUS エージェント %s 起動  %s:%d" % (VERSION, a.host, a.port))
    log("  ホスト=%s ユーザー=%s セッション=%d" % (socket.gethostname(),
                                                 os.environ.get("USERNAME"), sid.value))
    if sid.value == 0:
        log("  ！セッション0で動いている。画面は取れず GUI も触れない。"
            "ログオン中のデスクトップから起動し直すこと")
    if not a.allow:
        log("  ！--allow が無い。LAN の誰でもつなげる状態になっている")
    log("  任意コマンドの実行=%s  ファイルの範囲=%s"
        % ("閉じている" if a.no_exec else "開いている", a.root or "制限なし"))

    srv = ThreadingHTTPServer((a.host, a.port), Handler)
    srv.daemon_threads = True
    try:
        srv.serve_forever()
    except KeyboardInterrupt:
        log("停止")


if __name__ == "__main__":
    main()
