Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
7 changes: 5 additions & 2 deletions pyrit/prompt_target/http_target/http_target.py
Original file line number Diff line number Diff line change
Expand Up @@ -196,8 +196,11 @@ async def _send_prompt_to_target_async(self, *, normalized_conversation: list[Me
header_dict, http_body, url, http_method, http_version = self.parse_raw_http_request(http_request_w_prompt)
self._validate_destination(url)

if "Content-Length" in header_dict:
header_dict["Content-Length"] = str(len(http_body))
# Calculate and set content-length manually so it is available for request tracing
if isinstance(http_body, str) and http_body:
header_dict["content-length"] = str(len(http_body.encode("utf-8")))
elif isinstance(http_body, (bytes, bytearray)) and http_body:
header_dict["content-length"] = str(len(http_body))

http2_version = False
if http_version and "HTTP/2" in http_version:
Expand Down
71 changes: 66 additions & 5 deletions tests/unit/prompt_target/target/test_http_target.py
Original file line number Diff line number Diff line change
Expand Up @@ -84,7 +84,7 @@ async def test_send_prompt_async(mock_request, mock_http_target, mock_http_respo
mock_request.assert_called_with(
method="POST",
url="https://example.com/",
headers={"host": "example.com", "content-type": "application/json"},
headers={"host": "example.com", "content-type": "application/json", "content-length": "25"},
content='{"prompt": "test_prompt"}',
follow_redirects=True,
)
Expand Down Expand Up @@ -356,7 +356,7 @@ async def test_send_prompt_regex_parse_async(mock_request, mock_http_target):
mock_request.assert_called_with(
method="POST",
url="https://example.com/",
headers={"host": "example.com", "content-type": "application/json"},
headers={"host": "example.com", "content-type": "application/json", "content-length": "25"},
content='{"prompt": "test_prompt"}',
follow_redirects=True,
)
Expand Down Expand Up @@ -389,7 +389,7 @@ async def test_send_prompt_async_keeps_original_template(mock_request, mock_http
mock_request.assert_called_with(
method="POST",
url="https://example.com/",
headers={"host": "example.com", "content-type": "application/json"},
headers={"host": "example.com", "content-type": "application/json", "content-length": "25"},
content='{"prompt": "test_prompt"}',
follow_redirects=True,
)
Expand Down Expand Up @@ -417,14 +417,14 @@ async def test_send_prompt_async_keeps_original_template(mock_request, mock_http
mock_request.assert_any_call(
method="POST",
url="https://example.com/",
headers={"host": "example.com", "content-type": "application/json"},
headers={"host": "example.com", "content-type": "application/json", "content-length": "25"},
content='{"prompt": "test_prompt"}',
follow_redirects=True,
)
mock_request.assert_any_call(
method="POST",
url="https://example.com/",
headers={"host": "example.com", "content-type": "application/json"},
headers={"host": "example.com", "content-type": "application/json", "content-length": "32"},
content='{"prompt": "second_test_prompt"}',
follow_redirects=True,
)
Expand Down Expand Up @@ -542,3 +542,64 @@ def test_http_target_init_with_client_only():
)
assert target._client is client
assert target.httpx_client_kwargs == {}


@patch("httpx.AsyncClient.request", new_callable=AsyncMock)
async def test_send_prompt_async_content_length_utf8_bytes(mock_request, patch_central_database):
# Setup HTTPTarget
request = "POST / HTTP/1.1\nHost: example.com\n\n{PROMPT}"
target = HTTPTarget(http_request=request)

# Send a prompt containing multi-byte characters (e.g., emojis or non-ASCII text)
# "你好" has 2 characters, but in UTF-8 it is 6 bytes.
message = Message(message_pieces=[MessagePiece(role="user", original_value="你好")])
mock_response = MagicMock()
mock_response.content = b"ok"
mock_request.return_value = mock_response

await target.send_prompt_async(message=message)

# Validate that the requested header dict updated Content-Length to the correct byte length
assert mock_request.call_args.kwargs["headers"]["content-length"] == "6"


@patch("httpx.AsyncClient.request", new_callable=AsyncMock)
async def test_send_prompt_async_content_length_recalculated(mock_request, patch_central_database):
# Setup HTTPTarget with an old content-length
request = "POST / HTTP/1.1\nHost: example.com\ncontent-length: 0\n\n{PROMPT}"
target = HTTPTarget(http_request=request)

# Send standard ASCII payload
message = Message(message_pieces=[MessagePiece(role="user", original_value="test")])
mock_response = MagicMock()
mock_response.content = b"ok"
mock_request.return_value = mock_response

await target.send_prompt_async(message=message)

# Validate that it correctly identified and updated the header
assert mock_request.call_args.kwargs["headers"]["content-length"] == "4"


@patch("httpx.AsyncClient.request", new_callable=AsyncMock)
async def test_send_prompt_async_content_length_raw_bytes(mock_request, patch_central_database):
target = HTTPTarget(http_request="POST / HTTP/1.1\nHost: example.com\n\n")
message = Message(message_pieces=[MessagePiece(role="user", original_value="test")])
mock_response = MagicMock()
mock_response.content = b"ok"
mock_request.return_value = mock_response

with patch.object(
target,
"parse_raw_http_request",
return_value=(
{"host": "example.com"},
b"raw_byte_data",
"https://example.com/",
"POST",
"HTTP/1.1",
),
):
await target.send_prompt_async(message=message)

assert mock_request.call_args.kwargs["headers"]["content-length"] == "13"