Skip to content
Merged
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
2 changes: 1 addition & 1 deletion pyrit/prompt_target/http_target/http_target.py
Original file line number Diff line number Diff line change
Expand Up @@ -229,7 +229,7 @@ async def _send_prompt_to_target_async(self, *, normalized_conversation: list[Me
follow_redirects=self.follow_redirects,
)

response_content = response.content
response_content = response.text

if self.callback_function:
response_content = self.callback_function(response=response)
Expand Down
32 changes: 19 additions & 13 deletions pyrit/prompt_target/http_target/http_target_callback_functions.py
Original file line number Diff line number Diff line change
Expand Up @@ -7,10 +7,10 @@
from collections.abc import Callable
from typing import Any

import requests
import httpx


def get_http_target_json_response_callback_function(key: str) -> Callable[[requests.Response], str]:
def get_http_target_json_response_callback_function(key: str) -> Callable[[httpx.Response], str]:
"""
Determine proper parsing response function for an HTTP Request.

Expand All @@ -24,7 +24,7 @@ def get_http_target_json_response_callback_function(key: str) -> Callable[[reque
Callable: proper output parsing response
"""

def parse_json_http_response(response: requests.Response) -> str:
def parse_json_http_response(response: httpx.Response) -> str:
"""
Parse JSON outputs.

Expand All @@ -43,7 +43,7 @@ def parse_json_http_response(response: requests.Response) -> str:

def get_http_target_regex_matching_callback_function(
key: str, url: str | None = None
) -> Callable[[requests.Response], str]:
) -> Callable[[httpx.Response], str]:
"""
Get a callback function that parses HTTP responses using regex matching.

Expand All @@ -55,7 +55,7 @@ def get_http_target_regex_matching_callback_function(
Callable: A function that parses responses using the provided regex pattern.
"""

def parse_using_regex(response: requests.Response) -> str:
def parse_using_regex(response: httpx.Response) -> str:
"""
Parse text outputs using regex.

Expand All @@ -67,13 +67,16 @@ def parse_using_regex(response: requests.Response) -> str:
Returns:
str: parsed output from response given a regex pattern to follow
"""
# Search the decoded text; str(response.content) would search the bytes repr (b'...'),
# where non-ASCII characters and newlines appear as escape sequences.
response_text = response.text
re_pattern = re.compile(key)
match = re.search(re_pattern, str(response.content))
match = re.search(re_pattern, response_text)
if match:
if url:
return url + match.group()
return match.group()
return str(response.content)
return response_text

return parse_using_regex

Expand All @@ -93,14 +96,17 @@ def _fetch_key(data: dict[str, Any], key: str) -> Any:
ValueError: If any path segment is missing, so a misconfigured key
surfaces immediately instead of silently degrading to "".
"""
pattern = re.compile(r"([a-zA-Z_]+)|\[(-?\d+)\]")
# A key segment is a quoted bracket lookup, or any run of characters other than
# ".", "[" and "]", so keys such as "output2" or "generated-text" are kept whole.
pattern = re.compile(r"\[\s*[\"']([^\"']+)[\"']\s*\]|([^.\[\]]+)|\[(-?\d+)\]")
keys = pattern.findall(key)
result: Any = data
for key_part, index_part in keys:
if key_part:
if not isinstance(result, dict) or key_part not in result:
raise ValueError(f"Key path {key!r} not found in HTTP JSON response: missing segment {key_part!r}.")
result = result[key_part]
for quoted_part, key_part, index_part in keys:
name = quoted_part or key_part
if name:
if not isinstance(result, dict) or name not in result:
raise ValueError(f"Key path {key!r} not found in HTTP JSON response: missing segment {name!r}.")
result = result[name]
elif index_part:
index = int(index_part)
if not isinstance(result, list) or not -len(result) <= index < len(result):
Expand Down
2 changes: 1 addition & 1 deletion pyrit/prompt_target/http_target/httpx_api_target.py
Original file line number Diff line number Diff line change
Expand Up @@ -203,7 +203,7 @@ async def _send_prompt_to_target_async(self, *, normalized_conversation: list[Me
logger.error(f"File not found: {self.file_path}. Exception: {e}")
raise

response_content = response.content
response_content = response.text

# If a callback function was set, let them parse the response
if self.callback_function:
Expand Down
19 changes: 15 additions & 4 deletions tests/unit/prompt_target/target/test_http_api_target.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
from pathlib import Path
from unittest.mock import MagicMock, patch

import httpx
import pytest

from pyrit.models import Message, MessagePiece, RequestTraceContext
Expand All @@ -24,8 +25,7 @@ async def test_send_prompt_async_file_upload(mock_request, patch_central_databas
message = Message(message_pieces=[message_piece])

# Mock a response simulating a file upload.
mock_response = MagicMock()
mock_response.content = b'{"message": "File uploaded successfully", "filename": "mock.pdf"}'
mock_response = httpx.Response(200, content=b'{"message": "File uploaded successfully", "filename": "mock.pdf"}')
mock_request.return_value = mock_response

# Create HTTPXAPITarget without passing a transport.
Expand Down Expand Up @@ -87,8 +87,7 @@ async def test_send_prompt_async_no_file(mock_request, patch_central_database):
message = Message(message_pieces=[message_piece])

# Mock a response simulating a standard API (non-file).
mock_response = MagicMock()
mock_response.content = b'{"status": "ok", "data": "Sample JSON response"}'
mock_response = httpx.Response(200, content=b'{"status": "ok", "data": "Sample JSON response"}')
mock_request.return_value = mock_response

target = HTTPXAPITarget(http_url="http://example.com/data/", method="POST", timeout=180)
Expand All @@ -104,6 +103,18 @@ async def test_send_prompt_async_no_file(mock_request, patch_central_database):
assert "Sample JSON response" in response_text


@patch("httpx.AsyncClient.request")
async def test_send_prompt_async_stores_decoded_text(mock_request, patch_central_database):
message = Message(message_pieces=[MessagePiece(role="user", original_value="mock", converted_value="hello")])
body = "R\u00e9ponse \u2014 \U0001f600"
mock_request.return_value = httpx.Response(200, content=body.encode("utf-8"))

target = HTTPXAPITarget(http_url="http://example.com/data/", method="POST", timeout=180)
response = await target.send_prompt_async(message=message)

assert response[0].get_value() == body


@patch("httpx.AsyncClient.request")
async def test_send_prompt_async_preserves_query_params_for_post(mock_request, patch_central_database):
message_piece = MessagePiece(role="user", original_value="mock", converted_value="non_existent_file.pdf")
Expand Down
15 changes: 13 additions & 2 deletions tests/unit/prompt_target/target/test_http_target.py
Original file line number Diff line number Diff line change
Expand Up @@ -218,6 +218,18 @@ async def test_send_prompt_async_allows_configured_internal_destination(mock_req
assert mock_request.call_args.kwargs["url"] == "https://10.0.0.8:8080/api/jobs"


@patch("httpx.AsyncClient.request", new_callable=AsyncMock)
async def test_send_prompt_async_without_callback_stores_decoded_text(mock_request, patch_central_database):
target = HTTPTarget(http_request="POST /api HTTP/1.1\nHost: example.com\n\n{PROMPT}")
message = Message(message_pieces=[MessagePiece(role="user", original_value="prompt")])
body = "Sure \u2014 here\u2019s the answer:\nStep 1: caf\u00e9"
mock_request.return_value = httpx.Response(200, content=body.encode("utf-8"))

response = await target.send_prompt_async(message=message)

assert response[0].get_value() == body


@patch("httpx.AsyncClient.request", new_callable=AsyncMock)
async def test_send_prompt_async_follows_redirects_when_enabled(mock_request, patch_central_database):
target = HTTPTarget(
Expand Down Expand Up @@ -334,8 +346,7 @@ async def test_send_prompt_regex_parse_async(mock_request, mock_http_target):
)
]

mock_response = MagicMock()
mock_response.content = b"<html><body>Match: 1234</body></html>"
mock_response = httpx.Response(200, content=b"<html><body>Match: 1234</body></html>")
mock_request.return_value = mock_response

response = await mock_http_target.send_prompt_async(message=message)
Expand Down
66 changes: 61 additions & 5 deletions tests/unit/prompt_target/target/test_http_target_parsing.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
from collections.abc import Callable
from unittest.mock import MagicMock

import httpx
import pytest

from pyrit.prompt_target.http_target.http_target import HTTPTarget
Expand Down Expand Up @@ -103,16 +104,14 @@ def test_parse_raw_http_request_preserves_body_trailing_whitespace(sqlite_instan


def test_parse_regex_response_no_match():
mock_response = MagicMock()
mock_response.content = b"<html><body>No match here</body></html>"
mock_response = httpx.Response(200, content=b"<html><body>No match here</body></html>")
parse_html_function = get_http_target_regex_matching_callback_function(key=r'no_results\/[^\s"]+')
result = parse_html_function(mock_response)
assert result == "b'<html><body>No match here</body></html>'"
assert result == "<html><body>No match here</body></html>"


def test_parse_regex_response_match():
mock_response = MagicMock()
mock_response.content = b"<html><body>Match: 1234</body></html>"
mock_response = httpx.Response(200, content=b"<html><body>Match: 1234</body></html>")
parse_html_response = get_http_target_regex_matching_callback_function(r"Match: (\d+)")
result = parse_html_response(mock_response)
assert result == "Match: 1234"
Expand Down Expand Up @@ -156,3 +155,60 @@ def test_parse_json_response_out_of_range_index_raises():
parse_json_response = get_http_target_json_response_callback_function(key="data[5]")
with pytest.raises(ValueError, match=r"data\[5\]"):
parse_json_response(mock_response)


def test_parse_regex_response_matches_decoded_text():
body = "Sure \u2014 here\u2019s the answer:\nStep 1: caf\u00e9"
mock_response = httpx.Response(200, content=body.encode("utf-8"))
parse_html_response = get_http_target_regex_matching_callback_function(r"Step 1: .*")
result = parse_html_response(mock_response)
assert result == "Step 1: caf\u00e9"


def test_parse_regex_response_no_match_returns_decoded_text():
body = "Na\u00efve r\u00e9sum\u00e9 \U0001f600"
mock_response = httpx.Response(200, content=body.encode("utf-8"))
parse_html_response = get_http_target_regex_matching_callback_function(r"no_match")
assert parse_html_response(mock_response) == body


@pytest.mark.parametrize(
("key", "expected"),
[
("output2", "digits"),
("generated-text", "hyphen"),
("data.v2_answer", "nested"),
("results[1].item-id", "list"),
],
)
def test_parse_json_response_keys_with_digits_and_hyphens(key: str, expected: str):
mock_response = httpx.Response(
200,
content=(
b'{"output2": "digits", "generated-text": "hyphen", "data": {"v2_answer": "nested"},'
b' "results": [{"item-id": "x"}, {"item-id": "list"}]}'
),
)
parse_json_response = get_http_target_json_response_callback_function(key=key)
assert parse_json_response(mock_response) == expected


@pytest.mark.parametrize(
"key",
[
'choices[0]["message"]["content"]',
"choices[0]['message']['content']",
'choices[0][ "message" ].content',
"choices[-1].message.content",
],
)
def test_parse_json_response_quoted_bracket_keys(key: str):
mock_response = httpx.Response(200, content=b'{"choices": [{"message": {"content": "hello"}}]}')
parse_json_response = get_http_target_json_response_callback_function(key=key)
assert parse_json_response(mock_response) == "hello"


def test_parse_json_response_quoted_key_may_contain_dots():
mock_response = httpx.Response(200, content=b'{"data": {"model.name": "gpt"}}')
parse_json_response = get_http_target_json_response_callback_function(key='data["model.name"]')
assert parse_json_response(mock_response) == "gpt"
Loading