[sans-io] fix serialization, fixup client replay

This commit is contained in:
Maximilian Hils 2020-12-05 20:21:48 +01:00
parent 011766b10d
commit d55055d022
6 changed files with 81 additions and 29 deletions

View file

@ -1,4 +1,5 @@
import asyncio
import time
import traceback
import typing
@ -14,7 +15,7 @@ from mitmproxy.net import server_spec
from mitmproxy.options import Options
from mitmproxy.proxy.protocol.http import HTTPMode
from mitmproxy.proxy2 import commands, events, layers, server
from mitmproxy.proxy2.context import Context, Server
from mitmproxy.proxy2.context import ConnectionState, Context, Server
from mitmproxy.proxy2.layer import CommandGenerator
from mitmproxy.utils import asyncio_utils
@ -33,7 +34,13 @@ class MockServer(layers.http.HttpConnection):
def _handle_event(self, event: events.Event) -> CommandGenerator[None]:
if isinstance(event, events.Start):
has_content = bool(self.flow.request.raw_content)
yield layers.http.ReceiveHttp(layers.http.RequestHeaders(1, self.flow.request, not has_content))
self.flow.request.timestamp_start = self.flow.request.timestamp_end = time.time()
yield layers.http.ReceiveHttp(layers.http.RequestHeaders(
1,
self.flow.request,
end_stream=not has_content,
replay_flow=self.flow,
))
if has_content:
yield layers.http.ReceiveHttp(layers.http.RequestData(1, self.flow.request.raw_content))
yield layers.http.ReceiveHttp(layers.http.RequestEndOfMessage(1))
@ -51,6 +58,7 @@ class MockServer(layers.http.HttpConnection):
class ReplayHandler(server.ConnectionHandler):
def __init__(self, flow: http.HTTPFlow, options: Options) -> None:
client = flow.client_conn.copy()
client.state = ConnectionState.OPEN
context = Context(client, options)
context.server = Server(
@ -121,7 +129,7 @@ class ClientPlayback:
self.inflight = None
def check(self, f: flow.Flow) -> typing.Optional[str]:
if f.live:
if f.live or f == self.inflight:
return "Can't replay live flow."
if f.intercepted:
return "Can't replay intercepted flow."

View file

@ -1,6 +1,6 @@
import uuid
import warnings
from enum import Flag, auto
from enum import Flag
from typing import List, Literal, Optional, Sequence, Tuple, Union
from mitmproxy import certs
@ -11,8 +11,8 @@ from mitmproxy.options import Options
class ConnectionState(Flag):
CLOSED = 0
CAN_READ = auto()
CAN_WRITE = auto()
CAN_READ = 1
CAN_WRITE = 2
OPEN = CAN_READ | CAN_WRITE
@ -107,8 +107,7 @@ class Client(Connection):
def get_state(self):
# Important: Retain full compatibility with old proxy core for now!
# This means we truncate some fields (in either direction),
# which needs to be undone once we drop the old implementation.
# This means we need to add all new fields to the old implementation.
return {
'address': self.peername,
'alpn_proto_negotiated': self.alpn,
@ -123,6 +122,14 @@ class Client(Connection):
'tls_established': self.tls_established,
'tls_extensions': [],
'tls_version': self.tls_version,
# only used in sans-io
'state': self.state.value,
'sockname': self.sockname,
'error': self.error,
'tls': self.tls,
'certificate_list': [x.get_state() for x in self.certificate_list] if self.certificate_list else None,
'alpn_offers': self.alpn_offers,
'cipher_list': self.cipher_list,
}
@classmethod
@ -136,7 +143,7 @@ class Client(Connection):
return client
def set_state(self, state):
self.peername = state["address"]
self.peername = tuple(state["address"])
self.alpn = state["alpn_proto_negotiated"]
self.cipher = state["cipher_name"]
self.certificate_list = [certs.Cert.from_state(state["clientcert"])] if state["clientcert"] else None
@ -146,6 +153,15 @@ class Client(Connection):
self.timestamp_start = state["timestamp_start"]
self.timestamp_tls_setup = state["timestamp_tls_setup"]
self.tls_version = state["tls_version"]
# only used in sans-io
self.state = ConnectionState(state["state"])
self.sockname = tuple(state["sockname"])
self.error = state["error"]
self.tls = state["tls"]
self.certificate_list = [certs.Cert.from_state(x) for x in state["certificate_list"]] if state[
"certificate_list"] else None
self.alpn_offers = state["alpn_offers"]
self.cipher_list = state["cipher_list"]
@property
def address(self):
@ -192,7 +208,15 @@ class Server(Connection):
'timestamp_tls_setup': self.timestamp_tls_setup,
'tls_established': self.tls_established,
'tls_version': self.tls_version,
'via': None
'via': None,
# only used in sans-io
'state': self.state.value,
'error': self.error,
'tls': self.tls,
'certificate_list': [x.get_state() for x in self.certificate_list] if self.certificate_list else None,
'alpn_offers': self.alpn_offers,
'cipher_list': self.cipher_list,
'via2': self.via,
}
@classmethod
@ -202,18 +226,26 @@ class Server(Connection):
return server
def set_state(self, state):
self.address = state["address"]
self.address = tuple(state["address"])
self.alpn = state["alpn_proto_negotiated"]
self.certificate_list = [certs.Cert.from_state(state["cert"])] if state["cert"] else None
self.id = state["id"]
self.peername = state["ip_address"]
self.peername = tuple(state["ip_address"])
self.sni = state["sni"]
self.sockname = state["source_address"]
self.sockname = tuple(state["source_address"])
self.timestamp_end = state["timestamp_end"]
self.timestamp_start = state["timestamp_start"]
self.timestamp_tcp_setup = state["timestamp_tcp_setup"]
self.timestamp_tls_setup = state["timestamp_tls_setup"]
self.tls_version = state["tls_version"]
self.state = ConnectionState(state["state"])
self.error = state["error"]
self.tls = state["tls"]
self.certificate_list = [certs.Cert.from_state(x) for x in state["certificate_list"]] if state[
"certificate_list"] else None
self.alpn_offers = state["alpn_offers"]
self.cipher_list = state["cipher_list"]
self.via = state["via2"]
@property
def ip_address(self) -> Address:

View file

@ -5,7 +5,6 @@ import collections
import textwrap
import typing
from abc import abstractmethod
from dataclasses import dataclass
from mitmproxy import log
from mitmproxy.proxy2 import commands, events
@ -45,8 +44,13 @@ class Layer:
self._paused_event_queue = collections.deque()
show_debug_output = (
"termlog_verbosity" in context.options and
log.log_tier(context.options.termlog_verbosity) >= log.log_tier("debug")
(
"termlog_verbosity" in context.options and
log.log_tier(context.options.termlog_verbosity) >= log.log_tier("debug")
) or (
"console_eventlog_verbosity" in context.options and
log.log_tier(context.options.console_eventlog_verbosity) >= log.log_tier("debug")
)
)
if show_debug_output:
self.debug = " " * len(context.layers)

View file

@ -132,11 +132,14 @@ class HttpStream(layer.Layer):
@expect(RequestHeaders)
def state_wait_for_request_headers(self, event: RequestHeaders) -> layer.CommandGenerator[None]:
self.stream_id = event.stream_id
# noinspection PyTypeChecker
self.flow = http.HTTPFlow(
self.context.client,
self.context.server
)
if not event.replay_flow:
self.flow = http.HTTPFlow(
self.context.client,
self.context.server
)
else:
self.flow = event.replay_flow
self.flow.request = event.request
if err := validate_request(self.mode, self.flow.request):

View file

@ -1,6 +1,8 @@
from dataclasses import dataclass
from typing import Optional
from mitmproxy import http
from mitmproxy.http import HTTPFlow
from ._base import HttpEvent
@ -13,6 +15,8 @@ class RequestHeaders(HttpEvent):
us to set END_STREAM on headers already (and some servers - Akamai - implicitly expect that).
In either case, this event will nonetheless be followed by RequestEndOfMessage.
"""
replay_flow: Optional[HTTPFlow] = None
"""If set, the current request headers belong to a replayed flow, which should be reused."""
@dataclass

View file

@ -355,18 +355,19 @@ class Http2Client(Http2Connection):
]
if event.request.authority:
pseudo_headers.append((b":authority", event.request.data.authority))
elif not event.request.is_http2:
host_header = event.request.headers.pop("host", None)
if host_header:
pseudo_headers.append((b":authority", host_header))
headers = pseudo_headers + list(event.request.headers.fields)
if not event.request.is_http2:
headers = normalize_h1_headers(headers, True)
if event.request.is_http2:
hdrs = list(event.request.headers.fields)
else:
headers = event.request.headers
if not event.request.authority and "host" in headers:
headers = headers.copy()
pseudo_headers.append((b":authority", headers.pop(b"host")))
hdrs = normalize_h1_headers(list(headers.fields), True)
self.h2_conn.send_headers(
event.stream_id,
headers,
pseudo_headers + hdrs,
end_stream=event.end_stream,
)
self.streams[event.stream_id] = StreamState.EXPECTING_HEADERS