commit d822883e01144db69b0c73f9c3dc013f65991dca
parent aa40f0306c534721e3fd0a7e043bb29863a4a0a8
Author: triesap <tyson@radroots.org>
Date: Tue, 22 Sep 2026 15:35:46 +0000
test: strict bounded scripted provider fixtures (C002A)
Diffstat:
4 files changed, 480 insertions(+), 206 deletions(-)
diff --git a/tests/jev_provider_helper.mojo b/tests/jev_provider_helper.mojo
@@ -1,5 +1,5 @@
import std.os
-from std.ffi import c_int, c_size_t, c_ssize_t, external_call
+from std.ffi import c_int, c_size_t, c_ssize_t, c_uint, external_call
from std.os import Pipe, Process
from std.sys._libc import close
@@ -7,6 +7,8 @@ from flare.net import SocketAddr
from flare.tcp import TcpListener, TcpStream
from flare.utils import usleep
+from strict_fixture import ConnectionReader, FramedRequest, bearer_token
+
def _dup2(oldfd: c_int, newfd: c_int) -> c_int:
return external_call["dup2", c_int](oldfd, newfd)
@@ -18,6 +20,11 @@ def _fork() -> c_int:
@always_inline
+def _alarm(seconds: c_uint) -> c_uint:
+ return external_call["alarm", c_uint](seconds)
+
+
+@always_inline
def _exit_child(code: c_int):
_ = external_call["_exit", c_int](code)
@@ -44,67 +51,6 @@ def _read_pipe_line(mut pipe: Pipe) raises -> String:
return output^
-def _read_request(mut stream: TcpStream) raises -> String:
- var bytes = List[UInt8]()
- var chunk = InlineArray[Byte, 4096](fill=0)
- var expected_total = -1
- while True:
- var n = stream.read(chunk.unsafe_ptr(), 4096)
- if n <= 0:
- break
- for index in range(Int(n)):
- bytes.append(chunk[index])
- var text = String(unsafe_from_utf8=bytes[:])
- var header_end = text.find("\r\n\r\n")
- if header_end >= 0 and expected_total < 0:
- var lowered = text.lower()
- var marker = lowered.find("content-length:")
- var content_length = 0
- if marker >= 0:
- var rest = String(text[byte = marker + 15 :])
- var line_end = rest.find("\r\n")
- var value = rest if line_end < 0 else String(
- rest[byte=0:line_end]
- )
- content_length = Int(String(String(value).strip()))
- expected_total = header_end + 4 + content_length
- if expected_total >= 0 and len(bytes) >= expected_total:
- break
- if len(bytes) == 0:
- return ""
- return String(unsafe_from_utf8=bytes[:])
-
-
-def _request_path(request: String) -> String:
- var line_end = request.find("\r\n")
- if line_end < 0:
- return ""
- var first_line = String(request[byte=0:line_end])
- var first_space = first_line.find(" ")
- if first_space < 0:
- return ""
- var rest = String(first_line[byte = first_space + 1 :])
- var second_space = rest.find(" ")
- if second_space < 0:
- return ""
- return String(rest[byte=0:second_space])
-
-
-def _header_value(request: String, name: String) -> String:
- var header_end = request.find("\r\n\r\n")
- var header_block = request if header_end < 0 else String(
- request[byte=0:header_end]
- )
- var lowered = header_block.lower()
- var marker = lowered.find(name.lower() + ":")
- if marker < 0:
- return ""
- var rest = String(header_block[byte = marker + name.byte_length() + 1 :])
- var line_end = rest.find("\r\n")
- var value = rest if line_end < 0 else String(rest[byte=0:line_end])
- return String(String(value).strip())
-
-
def _json_string(value: String) -> String:
return '"' + value.replace("\\", "\\\\").replace('"', '\\"') + '"'
@@ -121,7 +67,9 @@ def _analysis() -> String:
def _response(status: Int, body: String) -> String:
var reason = "OK"
- if status == 429:
+ if status == 401:
+ reason = "Unauthorized"
+ elif status == 429:
reason = "Too Many Requests"
elif status == 500:
reason = "Internal Server Error"
@@ -139,82 +87,81 @@ def _response(status: Int, body: String) -> String:
)
-def _send(mut stream: TcpStream, status: Int, body: String) raises:
- stream.write_all(Span[UInt8, _](_response(status, body).as_bytes()))
+def _send(mut reader: ConnectionReader, status: Int, body: String) raises:
+ reader.write_all(_response(status, body))
-def _send_raw(mut stream: TcpStream, response: String) raises:
- stream.write_all(Span[UInt8, _](response.as_bytes()))
+def _send_raw(mut reader: ConnectionReader, response: String) raises:
+ reader.write_all(response)
-def _handle(mut stream: TcpStream, mode: String) raises:
- var request = _read_request(stream)
- var path = _request_path(request)
+def _handle(
+ mut reader: ConnectionReader, mode: String, framed: FramedRequest
+) raises:
+ var path = framed.path
if mode == "echo_authorization":
- _send(
- stream,
- 200,
- '{"authorization":"'
- + _header_value(request, "authorization")
- + '"}',
- )
- stream.close()
+ var token = bearer_token(framed.headers_raw)
+ if token == "":
+ _send(reader, 401, '{"error":{"message":"missing authorization"}}')
+ else:
+ _send(reader, 200, '{"authorization":' + _json_string(token) + "}")
return
if path != "/v1/systemone":
- _send(stream, 404, '{"error":{"message":"not found"}}')
- stream.close()
+ _send(reader, 404, '{"error":{"message":"not found"}}')
return
if mode == "ok":
- _send(stream, 200, _analysis())
+ _send(reader, 200, _analysis())
elif mode == "rate_limit":
- _send(stream, 429, '{"error":{"message":"slow down"}}')
+ _send(reader, 429, '{"error":{"message":"slow down"}}')
elif mode == "server_error":
- _send(stream, 500, '{"error":{"message":"boom"}}')
+ _send(reader, 500, '{"error":{"message":"boom"}}')
elif mode == "overloaded":
- _send(stream, 529, '{"error":{"message":"overloaded"}}')
+ _send(reader, 529, '{"error":{"message":"overloaded"}}')
elif mode == "auth":
- _send(stream, 401, '{"error":{"message":"bad key"}}')
+ _send(reader, 401, '{"error":{"message":"bad key"}}')
elif mode == "malformed_json":
- _send(stream, 200, "not json")
+ _send(reader, 200, "not json")
elif mode == "model_mismatch":
- _send(stream, 200, _analysis().replace("jev-1.13.0", "jev-other"))
+ _send(reader, 200, _analysis().replace("jev-1.13.0", "jev-other"))
elif mode == "truncated":
var truncated = (
"HTTP/1.1 200 OK\r\ncontent-type: application/json\r\n"
"content-length: 999\r\nconnection: close\r\n\r\n"
'{"model":"jev'
)
- stream.write_all(Span[UInt8, _](truncated.as_bytes()))
+ reader.write_all(truncated)
elif mode == "redirect":
- var body = ""
- stream.write_all(
- Span[UInt8, _](
- (
- "HTTP/1.1 302 Found\r\nlocation:"
- " http://127.0.0.1:1/steal\r\ncontent-length:"
- " 0\r\nconnection: close\r\n\r\n"
- ).as_bytes()
- )
+ reader.write_all(
+ "HTTP/1.1 302 Found\r\nlocation:"
+ " http://127.0.0.1:1/steal\r\ncontent-length:"
+ " 0\r\nconnection: close\r\n\r\n"
)
elif mode == "slow":
usleep(2_000_000)
- _send(stream, 200, _analysis())
+ _send(reader, 200, _analysis())
else:
- _send(stream, 500, '{"error":{"message":"unsupported_mode"}}')
- stream.close()
+ _send(reader, 500, '{"error":{"message":"unsupported_mode"}}')
def _serve(port: Int, mode: String, requests: Int) raises:
var listener = TcpListener.bind(SocketAddr.localhost(UInt16(port)))
var actual_port = Int(listener.local_addr().port)
_write(1, "ready " + String(actual_port) + "\n")
- for _ in range(requests):
+ var request_count = 0
+ while request_count < requests:
var stream = listener.accept()
- try:
- _handle(stream, mode)
- except:
- pass
+ var reader = ConnectionReader(stream^)
+ while request_count < requests:
+ var framed = reader.read()
+ if not framed.ok:
+ if framed.error == "empty":
+ break
+ raise Error("strict fixture framing error: " + framed.error)
+ request_count += 1
+ _handle(reader, mode, framed)
+ if not framed.keep_alive:
+ break
listener.close()
@@ -271,6 +218,7 @@ def _spawn_jev_stub(
_exit_child(c_int(126))
_ = close(stdout_read_fd)
_ = close(stdout_write_fd)
+ _ = _alarm(c_uint(20))
try:
_serve(port, mode, requests)
_exit_child(c_int(0))
diff --git a/tests/max_local_process_helper.mojo b/tests/max_local_process_helper.mojo
@@ -1,4 +1,4 @@
-from std.ffi import c_int, c_size_t, c_ssize_t, external_call
+from std.ffi import c_int, c_size_t, c_ssize_t, c_uint, external_call
from std.os import Pipe, Process
from std.sys._libc import close
@@ -7,6 +7,8 @@ from flare.tcp import TcpListener
from flare.tcp import TcpStream
from flare.utils import usleep
+from strict_fixture import ConnectionReader, FramedRequest, bearer_token
+
def _dup2(oldfd: c_int, newfd: c_int) -> c_int:
return external_call["dup2", c_int](oldfd, newfd)
@@ -18,6 +20,16 @@ def _fork() -> c_int:
@always_inline
+def _kill(pid: c_int, sig: c_int) -> c_int:
+ return external_call["kill", c_int](pid, sig)
+
+
+@always_inline
+def _alarm(seconds: c_uint) -> c_uint:
+ return external_call["alarm", c_uint](seconds)
+
+
+@always_inline
def _exit_child(code: c_int):
_ = external_call["_exit", c_int](code)
@@ -44,64 +56,6 @@ def _read_pipe_line(mut pipe: Pipe) raises -> String:
return output^
-comptime MAX_TEST_REQUEST_BYTES = 1048576
-
-
-def _read_request(mut stream: TcpStream) raises -> String:
- var bytes = List[UInt8]()
- var chunk = InlineArray[Byte, 4096](fill=0)
- var expected_total = -1
- while True:
- var n = stream.read(chunk.unsafe_ptr(), 4096)
- if n <= 0:
- break
- for index in range(Int(n)):
- bytes.append(chunk[index])
- if len(bytes) > MAX_TEST_REQUEST_BYTES:
- break
- var text = String(unsafe_from_utf8=bytes[:])
- var header_end = text.find("\r\n\r\n")
- if header_end >= 0 and expected_total < 0:
- var lowered = text.lower()
- var marker = lowered.find("content-length:")
- var content_length = 0
- if marker >= 0:
- var rest = String(text[byte = marker + 15 :])
- var line_end = rest.find("\r\n")
- var value = rest if line_end < 0 else String(
- rest[byte=0:line_end]
- )
- content_length = Int(String(String(value).strip()))
- expected_total = header_end + 4 + content_length
- if expected_total >= 0 and len(bytes) >= expected_total:
- break
- if len(bytes) == 0:
- return ""
- return String(unsafe_from_utf8=bytes[:])
-
-
-def _request_body(request: String) -> String:
- var header_end = request.find("\r\n\r\n")
- if header_end < 0:
- return ""
- return String(request[byte = header_end + 4 :])
-
-
-def _request_path(request: String) -> String:
- var line_end = request.find("\r\n")
- if line_end < 0:
- return ""
- var first_line = String(request[byte=0:line_end])
- var first_space = first_line.find(" ")
- if first_space < 0:
- return ""
- var rest = String(first_line[byte = first_space + 1 :])
- var second_space = rest.find(" ")
- if second_space < 0:
- return ""
- return String(rest[byte=0:second_space])
-
-
def _json_string(value: String) -> String:
return '"' + value.replace("\\", "\\\\").replace('"', '\\"') + '"'
@@ -128,7 +82,9 @@ def _chat_completion(body: String) -> String:
def _response(status: Int, body: String) -> String:
var reason = "OK"
- if status == 404:
+ if status == 401:
+ reason = "Unauthorized"
+ elif status == 404:
reason = "Not Found"
elif status == 500:
reason = "Internal Server Error"
@@ -146,36 +102,35 @@ def _response(status: Int, body: String) -> String:
)
-def _send(mut stream: TcpStream, status: Int, body: String) raises:
- var response = _response(status, body)
- stream.write_all(Span[UInt8, _](response.as_bytes()))
+def _send(mut reader: ConnectionReader, status: Int, body: String) raises:
+ reader.write_all(_response(status, body))
-def _send_raw(mut stream: TcpStream, response: String) raises:
- stream.write_all(Span[UInt8, _](response.as_bytes()))
+def _send_raw(mut reader: ConnectionReader, response: String) raises:
+ reader.write_all(response)
-def _handle_health(mut stream: TcpStream, mode: String) raises:
+def _handle_health(mut reader: ConnectionReader, mode: String) raises:
if mode == "health_non_2xx":
- _send(stream, 503, '{"status":"unavailable"}')
+ _send(reader, 503, '{"status":"unavailable"}')
elif mode == "health_timeout":
usleep(1_000_000)
elif mode == "health_malformed_http":
- _send_raw(stream, "not an http response\r\n\r\n")
+ _send_raw(reader, "not an http response\r\n\r\n")
elif mode == "query_rewrite_remaining_deadline_timeout":
usleep(200_000)
- _send(stream, 200, '{"status":"ok"}')
+ _send(reader, 200, '{"status":"ok"}')
else:
- _send(stream, 200, '{"status":"ok"}')
+ _send(reader, 200, '{"status":"ok"}')
-def _handle_chat_completions(mut stream: TcpStream, mode: String) raises:
+def _handle_chat_completions(mut reader: ConnectionReader, mode: String) raises:
if mode == "query_rewrite_ok":
- _send(stream, 200, _chat_completion(_query_rewrite_analysis()))
+ _send(reader, 200, _chat_completion(_query_rewrite_analysis()))
elif mode == "query_rewrite_non_2xx":
- _send(stream, 503, '{"error":{"message":"provider unavailable"}}')
+ _send(reader, 503, '{"error":{"message":"provider unavailable"}}')
elif mode == "query_rewrite_invalid_json":
- _send(stream, 200, '{"choices":[{"message":{"content":"not json"}}]}')
+ _send(reader, 200, '{"choices":[{"message":{"content":"not json"}}]}')
elif mode == "query_rewrite_schema_invalid":
var body = (
'{"original_text":"local apples pickup weekend",'
@@ -189,60 +144,93 @@ def _handle_chat_completions(mut stream: TcpStream, mode: String) raises:
'"time_window":"weekend"'
"}}"
)
- _send(stream, 200, _chat_completion(body))
+ _send(reader, 200, _chat_completion(body))
elif mode == "query_rewrite_top_level_string":
- _send(stream, 200, '"not object"')
+ _send(reader, 200, '"not object"')
elif mode == "query_rewrite_top_level_array":
- _send(stream, 200, "[]")
+ _send(reader, 200, "[]")
elif mode == "query_rewrite_top_level_null":
- _send(stream, 200, "null")
+ _send(reader, 200, "null")
elif mode == "query_rewrite_empty_choices":
- _send(stream, 200, '{"choices":[]}')
+ _send(reader, 200, '{"choices":[]}')
elif mode == "query_rewrite_missing_content":
- _send(stream, 200, '{"choices":[{"message":{}}]}')
+ _send(reader, 200, '{"choices":[{"message":{}}]}')
elif mode == "query_rewrite_error_payload":
- _send(stream, 200, '{"error":{"message":"provider refusal"}}')
+ _send(reader, 200, '{"error":{"message":"provider refusal"}}')
elif mode == "query_rewrite_timeout":
usleep(2_000_000)
- _send(stream, 200, _chat_completion(_query_rewrite_analysis()))
+ _send(reader, 200, _chat_completion(_query_rewrite_analysis()))
elif mode == "query_rewrite_remaining_deadline_timeout":
usleep(400_000)
- _send(stream, 200, _chat_completion(_query_rewrite_analysis()))
+ _send(reader, 200, _chat_completion(_query_rewrite_analysis()))
elif mode == "query_rewrite_malformed_http":
- _send_raw(stream, "not an http response\r\n\r\n")
+ _send_raw(reader, "not an http response\r\n\r\n")
else:
- _send(stream, 500, '{"error":"unsupported_mode"}')
+ _send(reader, 500, '{"error":"unsupported_mode"}')
-def _handle_request(mut stream: TcpStream, mode: String, index: Int) raises:
- var request = _read_request(stream)
- var path = _request_path(request)
+def _handle_framed(
+ mut reader: ConnectionReader,
+ mode: String,
+ request_index: Int,
+ connection_index: Int,
+ framed: FramedRequest,
+) raises:
+ var path = framed.path
if mode == "echo_body_bytes":
_send(
- stream,
+ reader,
200,
- '{"received_bytes":'
- + String(_request_body(request).byte_length())
- + "}",
+ '{"received_bytes":' + String(framed.body.byte_length()) + "}",
)
elif mode == "count_requests":
- _send(stream, 200, '{"request_index":' + String(index + 1) + "}")
+ _send(
+ reader,
+ 200,
+ '{"request_index":'
+ + String(request_index)
+ + ',"connection_index":'
+ + String(connection_index)
+ + "}",
+ )
+ elif mode == "echo_authorization":
+ var token = bearer_token(framed.headers_raw)
+ if token == "":
+ _send(reader, 401, '{"error":"missing_authorization"}')
+ else:
+ _send(reader, 200, '{"authorization":' + _json_string(token) + "}")
+ elif mode == "stall":
+ usleep(30_000_000)
elif path == "/health":
- _handle_health(stream, mode)
+ _handle_health(reader, mode)
elif path == "/v1/chat/completions":
- _handle_chat_completions(stream, mode)
+ _handle_chat_completions(reader, mode)
else:
- _send(stream, 404, '{"error":"not_found"}')
- stream.close()
+ _send(reader, 404, '{"error":"not_found"}')
def _serve_max_local_stub(port: Int, mode: String, requests: Int) raises:
var listener = TcpListener.bind(SocketAddr.localhost(UInt16(port)))
var actual_port = Int(listener.local_addr().port)
_write(1, "ready " + String(actual_port) + "\n")
- for request_index in range(requests):
+ var request_count = 0
+ var connection_count = 0
+ while request_count < requests:
var stream = listener.accept()
- _handle_request(stream, mode, request_index)
+ connection_count += 1
+ var reader = ConnectionReader(stream^)
+ while request_count < requests:
+ var framed = reader.read()
+ if not framed.ok:
+ if framed.error == "empty":
+ break
+ raise Error("strict fixture framing error: " + framed.error)
+ request_count += 1
+ _handle_framed(
+ reader, mode, request_count, connection_count, framed
+ )
+ if not framed.keep_alive:
+ break
listener.close()
@@ -260,6 +248,11 @@ struct SpawnedMaxLocalStub(Movable):
if not status.exit_code or status.exit_code.value() != 0:
raise Error("max_local stub exited unexpectedly")
+ def terminate(mut self) raises:
+ _ = _kill(c_int(self.pid), c_int(15))
+ var process = Process(self.pid)
+ _ = process.wait()
+
def reserve_loopback_port() raises -> Int:
var listener = TcpListener.bind(SocketAddr.localhost(0))
@@ -284,6 +277,7 @@ def spawn_max_local_stub(
_exit_child(c_int(126))
_ = close(stdout_read_fd)
_ = close(stdout_write_fd)
+ _ = _alarm(c_uint(20))
try:
_serve_max_local_stub(port, mode, requests)
_exit_child(c_int(0))
diff --git a/tests/strict_fixture.mojo b/tests/strict_fixture.mojo
@@ -0,0 +1,189 @@
+"""Strict bounded HTTP request framing for local test fixtures (ADR-0012 D29).
+
+Bounded byte-oriented accumulation and header parsing *before* UTF-8 decoding,
+validated nonnegative ``Content-Length``, exact header-name matching,
+duplicate/conflicting length rejection, premature-EOF rejection, bounded
+unsupported-transfer rejection and surplus retention for persistent
+connections.
+"""
+
+from std.collections import List
+
+from flare.tcp import TcpStream
+
+comptime STRICT_MAX_HEADER_BYTES = 65536
+comptime STRICT_MAX_BODY_BYTES = 1048576
+
+
+struct FramedRequest(Movable):
+ var ok: Bool
+ var error: String
+ var method: String
+ var path: String
+ var headers_raw: String
+ var body: String
+ var keep_alive: Bool
+ var total_bytes: Int
+
+ def __init__(out self):
+ self.ok = False
+ self.error = ""
+ self.method = ""
+ self.path = ""
+ self.headers_raw = ""
+ self.body = ""
+ self.keep_alive = False
+ self.total_bytes = 0
+
+
+def _bytes_string(bytes: List[UInt8], stop: Int) raises -> String:
+ return String(from_utf8=Span(ptr=bytes.unsafe_ptr(), length=stop))
+
+
+struct ConnectionReader(Movable):
+ var _stream: TcpStream
+ var _buffer: List[UInt8]
+ var _eof: Bool
+
+ def __init__(out self, var stream: TcpStream):
+ self._stream = stream^
+ self._buffer = List[UInt8]()
+ self._eof = False
+
+ def _read_more(mut self) raises -> Int:
+ var chunk = InlineArray[Byte, 4096](fill=0)
+ var n = self._stream.read(chunk.unsafe_ptr(), 4096)
+ if n <= 0:
+ self._eof = True
+ return 0
+ for index in range(Int(n)):
+ self._buffer.append(UInt8(Int(chunk[index])))
+ return Int(n)
+
+ def write_all(mut self, text: String) raises:
+ self._stream.write_all(Span[UInt8, _](text.as_bytes()))
+
+ def _header_end(self) -> Int:
+ var i = 0
+ while i + 3 < len(self._buffer):
+ if (
+ self._buffer[i] == 13
+ and self._buffer[i + 1] == 10
+ and self._buffer[i + 2] == 13
+ and self._buffer[i + 3] == 10
+ ):
+ return i
+ i += 1
+ return -1
+
+ def read(mut self) raises -> FramedRequest:
+ var outcome = FramedRequest()
+ var header_end = self._header_end()
+ while header_end < 0:
+ if self._eof:
+ outcome.error = (
+ "empty" if len(self._buffer)
+ == 0 else "missing_header_terminator"
+ )
+ return outcome^
+ if len(self._buffer) > STRICT_MAX_HEADER_BYTES:
+ outcome.error = "header_too_large"
+ return outcome^
+ _ = self._read_more()
+ header_end = self._header_end()
+ if header_end > STRICT_MAX_HEADER_BYTES:
+ outcome.error = "header_too_large"
+ return outcome^
+
+ var header_text = _bytes_string(self._buffer, header_end)
+ var lines = header_text.split("\r\n")
+ if len(lines) < 1:
+ outcome.error = "malformed_request_line"
+ return outcome^
+ var request_line = String(lines[0])
+ var parts = request_line.split(" ")
+ if len(parts) != 3:
+ outcome.error = "malformed_request_line"
+ return outcome^
+ outcome.method = String(parts[0])
+ outcome.path = String(parts[1])
+
+ var content_length = -1
+ var have_length = False
+ var have_transfer = False
+ for i in range(1, len(lines)):
+ var line = String(lines[i])
+ if line.byte_length() == 0:
+ continue
+ var colon = line.find(":")
+ if colon < 0:
+ outcome.error = "malformed_header"
+ return outcome^
+ var name = String(line[byte=0:colon]).lower().strip()
+ var value = String(line[byte = colon + 1 :]).strip()
+ if name == "content-length":
+ if have_length:
+ outcome.error = "duplicate_content_length"
+ return outcome^
+ var parsed = Int(value)
+ if parsed < 0:
+ outcome.error = "negative_content_length"
+ return outcome^
+ content_length = parsed
+ have_length = True
+ elif name == "transfer-encoding":
+ have_transfer = True
+ if value.lower() != "identity":
+ outcome.error = "unsupported_transfer_encoding"
+ return outcome^
+ elif name == "connection" and value.lower() == "keep-alive":
+ outcome.keep_alive = True
+
+ if have_transfer and have_length:
+ outcome.error = "conflicting_framing"
+ return outcome^
+ if not have_length:
+ content_length = 0
+ if content_length > STRICT_MAX_BODY_BYTES:
+ outcome.error = "body_too_large"
+ return outcome^
+
+ var expected_total = header_end + 4 + content_length
+ while len(self._buffer) < expected_total:
+ if self._eof:
+ outcome.error = "premature_eof"
+ return outcome^
+ _ = self._read_more()
+
+ outcome.headers_raw = header_text
+ outcome.body = String(
+ _bytes_string(self._buffer, expected_total)[byte = header_end + 4 :]
+ )
+ var surplus = List[UInt8]()
+ for index in range(expected_total, len(self._buffer)):
+ surplus.append(self._buffer[index])
+ self._buffer = surplus^
+ outcome.ok = True
+ outcome.total_bytes = expected_total
+ return outcome^
+
+
+def header_value(headers_raw: String, name: String) -> String:
+ """Return the value of an exactly named header, or "" when absent."""
+ var lines = headers_raw.split("\r\n")
+ var wanted = name.lower()
+ for i in range(1, len(lines)):
+ var line = String(lines[i])
+ var colon = line.find(":")
+ if colon < 0:
+ continue
+ if String(line[byte=0:colon]).lower().strip() == wanted:
+ return String(String(line[byte = colon + 1 :]).strip())
+ return ""
+
+
+def bearer_token(headers_raw: String) -> String:
+ var value = header_value(headers_raw, "authorization")
+ if not value.startswith("Bearer "):
+ return ""
+ return String(value[byte=7:])
diff --git a/tests/test_provider_helpers.mojo b/tests/test_provider_helpers.mojo
@@ -3,7 +3,10 @@ from std.testing import TestSuite, assert_true
from flare.net import SocketAddr
from flare.tcp import TcpStream
-from max_local_process_helper import spawn_max_local_stub
+from max_local_process_helper import (
+ SpawnedMaxLocalStub,
+ spawn_max_local_stub,
+)
from jev_provider_helper import spawn_jev_stub_auto
@@ -82,9 +85,149 @@ def test_jev_stub_observes_bearer_sentinel_at_intended_origin() raises:
var response = _client_request(
started.port, "/v1/systemone", "{}", "Bearer hyf-sentinel-token"
)
- assert_true(response.find("Bearer hyf-sentinel-token") >= 0)
+ assert_true(response.find("hyf-sentinel-token") >= 0)
started.stub.wait()
+def _client_send_raw(port: Int, data: String) raises -> String:
+ var client = TcpStream.connect(SocketAddr.localhost(UInt16(port)))
+ client.write_all(Span[UInt8, _](data.as_bytes()))
+ var buffer = InlineArray[Byte, 4096](fill=0)
+ var response = String("")
+ while True:
+ var n = client.read(buffer.unsafe_ptr(), 4096)
+ if n <= 0:
+ break
+ response += String(
+ unsafe_from_utf8=Span(ptr=buffer.unsafe_ptr(), length=Int(n))
+ )
+ client.close()
+ return response^
+
+
+def _client_two_on_one(port: Int, path: String) raises -> String:
+ var client = TcpStream.connect(SocketAddr.localhost(UInt16(port)))
+ var body = "{}"
+ var first = (
+ "POST "
+ + path
+ + " HTTP/1.1\r\nhost: 127.0.0.1\r\ncontent-type:"
+ " application/json\r\ncontent-length: "
+ + String(body.byte_length())
+ + "\r\nconnection: keep-alive\r\n\r\n"
+ + body
+ )
+ var second = (
+ "POST "
+ + path
+ + " HTTP/1.1\r\nhost: 127.0.0.1\r\ncontent-type:"
+ " application/json\r\ncontent-length: "
+ + String(body.byte_length())
+ + "\r\nconnection: close\r\n\r\n"
+ + body
+ )
+ client.write_all(Span[UInt8, _](first.as_bytes()))
+ client.write_all(Span[UInt8, _](second.as_bytes()))
+ var buffer = InlineArray[Byte, 4096](fill=0)
+ var response = String("")
+ while True:
+ var n = client.read(buffer.unsafe_ptr(), 4096)
+ if n <= 0:
+ break
+ response += String(
+ unsafe_from_utf8=Span(ptr=buffer.unsafe_ptr(), length=Int(n))
+ )
+ client.close()
+ return response^
+
+
+def _stub_failed(mut stub: SpawnedMaxLocalStub) raises -> Bool:
+ try:
+ stub.wait()
+ return False
+ except:
+ return True
+
+
+def _client_send_only(port: Int, data: String) raises:
+ var client = TcpStream.connect(SocketAddr.localhost(UInt16(port)))
+ client.write_all(Span[UInt8, _](data.as_bytes()))
+ client.close()
+
+
+def test_max_local_strict_rejects_truncated_frame() raises:
+ var stub = spawn_max_local_stub(0, "echo_body_bytes", 1)
+ _client_send_only(
+ stub.port,
+ (
+ "POST /v1/chat/completions HTTP/1.1\r\nhost: 127.0.0.1\r\n"
+ "content-length: 100\r\nconnection: close\r\n\r\nshort"
+ ),
+ )
+ assert_true(_stub_failed(stub))
+
+
+def test_max_local_strict_rejects_duplicate_content_length() raises:
+ var stub = spawn_max_local_stub(0, "echo_body_bytes", 1)
+ _client_send_only(
+ stub.port,
+ (
+ "POST /v1/chat/completions HTTP/1.1\r\nhost: 127.0.0.1\r\n"
+ "content-length: 2\r\ncontent-length: 2\r\n"
+ "connection: close\r\n\r\n{}"
+ ),
+ )
+ assert_true(_stub_failed(stub))
+
+
+def test_max_local_strict_rejects_oversize_body() raises:
+ var stub = spawn_max_local_stub(0, "echo_body_bytes", 1)
+ _client_send_only(
+ stub.port,
+ (
+ "POST /v1/chat/completions HTTP/1.1\r\nhost: 127.0.0.1\r\n"
+ "content-length: 2097152\r\nconnection: close\r\n\r\n"
+ ),
+ )
+ assert_true(_stub_failed(stub))
+
+
+def test_max_local_strict_persistent_exchanges_distinct_counters() raises:
+ var stub = spawn_max_local_stub(0, "count_requests", 2)
+ var response = _client_two_on_one(stub.port, "/v1/chat/completions")
+ assert_true(response.find('"request_index":1') >= 0)
+ assert_true(response.find('"request_index":2') >= 0)
+ assert_true(response.find('"connection_index":1') >= 0)
+ assert_true(response.find('"connection_index":2') < 0)
+ stub.wait()
+
+
+def test_jev_strict_rejects_x_authorization() raises:
+ var started = spawn_jev_stub_auto("echo_authorization", 1)
+ var response = _client_send_raw(
+ started.port,
+ (
+ "POST /v1/systemone HTTP/1.1\r\nhost: 127.0.0.1\r\nx-authorization:"
+ " Bearer spoof\r\ncontent-type: application/json\r\ncontent-length:"
+ " 2\r\nconnection: close\r\n\r\n{}"
+ ),
+ )
+ assert_true(response.find("401") >= 0)
+ assert_true(response.find("spoof") < 0)
+ started.stub.wait()
+
+
+def test_max_local_stub_stalled_child_is_reaped() raises:
+ var stub = spawn_max_local_stub(0, "stall", 1)
+ _client_send_only(
+ stub.port,
+ (
+ "POST /v1/chat/completions HTTP/1.1\r\nhost: 127.0.0.1\r\n"
+ "content-length: 2\r\nconnection: close\r\n\r\n{}"
+ ),
+ )
+ stub.terminate()
+
+
def main() raises:
TestSuite.discover_tests[__functions_in_module()]().run()