|
6 | 6 | from urllib.parse import unquote |
7 | 7 | from http.server import HTTPServer, BaseHTTPRequestHandler |
8 | 8 | # local |
9 | | -from ..template import get_html_template, get_svg_template |
10 | | -from ..crypto import NullEncryptor |
11 | | -from ..page_builder import PageBuilder, Compression, Encoding, DEFAULT_TEMPLATE_FILE, DEFAULT_SVG_FILE |
| 9 | +from . import Subcommand |
| 10 | +from .template import get_initial_page_contents |
| 11 | +from .encryption import get_encryptor |
| 12 | +from .output import get_compression_list, get_encoding_list |
| 13 | +from ..template import get_html_template, get_svg_template, DEFAULT_HTML_TEMPLATE_PATH |
| 14 | +from ..page_builder import PageBuilder |
12 | 15 | from ..static_js import JS_DOWNLOAD, JS_DOWNLOAD_SVG |
13 | 16 |
|
14 | 17 |
|
15 | | -def register_server_argument_parser(ap: ArgumentParser): |
16 | | - ap.add_argument("-b", "--bind", default="0.0.0.0", help="IP address to bind to (default: 0.0.0.0)") |
17 | | - ap.add_argument("-p", "--port", type=int, default=8000, help="port to bind to (default: 8000)") |
| 18 | +def register_server_argument_parser(ap: ArgumentParser, subcommand: Subcommand): |
| 19 | + if subcommand == Subcommand.SERVE: |
| 20 | + ap_server = ap.add_argument_group("Server options") |
| 21 | + ap_server.add_argument("-b", "--bind", default="0.0.0.0", help="IP address to bind to (default: 0.0.0.0)") |
| 22 | + ap_server.add_argument("port", nargs="?", type=int, default=8000, help="port to bind to (default: 8000)") |
18 | 23 |
|
19 | 24 | class HTMLSmugglingServer(HTTPServer): |
20 | | - def __init__(self, server_address, RequestHandlerClass): |
| 25 | + def __init__(self, server_address, RequestHandlerClass, args): |
21 | 26 | super().__init__(server_address, RequestHandlerClass) |
22 | | - self.html_template = get_html_template(DEFAULT_TEMPLATE_FILE, "File Download", "File download link should be shown immediately") |
23 | | - self.svg_template = get_svg_template(DEFAULT_SVG_FILE) |
24 | | - self.encryptor = NullEncryptor() |
| 27 | + |
| 28 | + try: |
| 29 | + self.svg_template = get_svg_template(args.svg) |
| 30 | + except: |
| 31 | + raise Exception(f"Failed to load SVG file '{args.svg}'. Try specifying a different file with the --svg option") |
| 32 | + |
| 33 | + template_file = args.template or DEFAULT_HTML_TEMPLATE_PATH |
| 34 | + initial_page_contents = get_initial_page_contents(args) |
| 35 | + try: |
| 36 | + self.html_template = get_html_template(template_file, args.title, initial_page_contents) |
| 37 | + except: |
| 38 | + raise Exception(f"Failed to load template file '{template_file}'. Try specifying a different file with the --template option") |
| 39 | + |
| 40 | + self.compression_list = get_compression_list(args) |
| 41 | + self.encoding_list = get_encoding_list(args) |
| 42 | + self.obscure_action = args.obscure_action |
| 43 | + self.insert_debug_statements = args.console_log |
| 44 | + |
| 45 | + self.encryptor = get_encryptor(args) |
| 46 | + |
| 47 | + # self.html_template = get_html_template(DEFAULT_TEMPLATE_FILE, "File Download", "File download link should be shown immediately") |
| 48 | + # self.svg_template = get_svg_template(DEFAULT_SVG_FILE) |
| 49 | + # self.encryptor = NullEncryptor() |
25 | 50 |
|
26 | 51 |
|
27 | 52 | class HTMLSmugglingRequestHandler(BaseHTTPRequestHandler): |
@@ -60,38 +85,38 @@ def serve_html(self, path): |
60 | 85 | with open(path, "rb") as f: |
61 | 86 | file_contents = f.read() |
62 | 87 |
|
63 | | - html_page_builder = PageBuilder( |
64 | | - self.server.html_template, |
65 | | - JS_DOWNLOAD.replace("{{NAME}}", file_name), |
66 | | - self.server.encryptor, |
67 | | - compression_list = [Compression.NONE], |
68 | | - encoding_list = [Encoding.BASE64], |
69 | | - ) |
70 | | - html_str = html_page_builder.build_page(file_contents) |
| 88 | + file_contents = self.build_page(file_name, file_contents, self.server.html_template, JS_DOWNLOAD, False) |
71 | 89 |
|
72 | 90 | self.send_response(200) |
73 | 91 | self.send_header("Content-type", "text/html; charset=utf-8") |
74 | 92 | self.end_headers() |
75 | | - self.wfile.write(html_str.encode()) |
| 93 | + self.wfile.write(file_contents) |
76 | 94 |
|
77 | 95 | def serve_svg(self, path): |
78 | 96 | file_name = os.path.basename(path) |
79 | 97 | with open(path, "rb") as f: |
80 | 98 | file_contents = f.read() |
81 | 99 |
|
| 100 | + file_contents = self.build_page(file_name, file_contents, self.server.svg_template, JS_DOWNLOAD_SVG, True) |
| 101 | + |
| 102 | + self.send_response(200) |
| 103 | + self.send_header("Content-type", "image/svg+xml") |
| 104 | + self.end_headers() |
| 105 | + self.wfile.write(file_contents) |
| 106 | + |
| 107 | + def build_page(self, file_name: str, file_contents: bytes, template: str, js_payload: str, is_svg: bool) -> bytes: |
82 | 108 | html_page_builder = PageBuilder( |
83 | | - self.server.svg_template, |
84 | | - JS_DOWNLOAD_SVG.replace("{{NAME}}", file_name), |
| 109 | + template, |
| 110 | + js_payload.replace("{{NAME}}", file_name), |
85 | 111 | self.server.encryptor, |
86 | | - compression_list = [Compression.NONE], |
87 | | - encoding_list = [Encoding.BASE64], |
| 112 | + obscure_action=self.server.obscure_action, |
| 113 | + encode_library_as_base64=False, |
| 114 | + insert_debug_statements=self.server.insert_debug_statements, |
| 115 | + compression_list = self.server.compression_list, |
| 116 | + encoding_list = self.server.encoding_list, |
88 | 117 | ) |
89 | 118 | html_str = html_page_builder.build_page(file_contents) |
90 | | - |
91 | | - self.send_response(200) |
92 | | - self.send_header("Content-type", "text/html; charset=utf-8") |
93 | | - self.end_headers() |
94 | | - self.wfile.write(html_str.encode()) |
| 119 | + return html_str.encode() |
95 | 120 |
|
96 | 121 | def list_directory(self, path): |
97 | 122 | try: |
@@ -133,10 +158,10 @@ def translate_path(self, path): |
133 | 158 | return base_path |
134 | 159 |
|
135 | 160 |
|
136 | | -def start_server(bind_ip: str, bind_port: int): |
| 161 | +def start_server(bind_ip: str, bind_port: int, args): |
137 | 162 | try: |
138 | 163 | server_address = (bind_ip, bind_port) |
139 | | - httpd = HTMLSmugglingServer(server_address, HTMLSmugglingRequestHandler) |
| 164 | + httpd = HTMLSmugglingServer(server_address, HTMLSmugglingRequestHandler, args) |
140 | 165 | print(f"Serving at http://{bind_ip}:{bind_port}") |
141 | 166 | httpd.serve_forever() |
142 | 167 | except KeyboardInterrupt: |
|
0 commit comments