Skip to content
Closed
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
14 changes: 14 additions & 0 deletions garak/generators/base.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,11 +42,25 @@ class Generator(Configurable):
# legal element for str list `modality['in']`: 'text', 'image', 'audio', 'video', '3d'
# refer to Table 1 in https://arxiv.org/abs/2401.13601
modality: dict = {"in": {"text"}, "out": {"text"}}
audio_formats: set[str] = set()
image_formats: set[str] = set()
video_formats: set[str] = set()
file_formats: set[str] = set()

supports_multiple_generations = (
False # can more than one generation be extracted per request?
)

@classmethod
def supported_formats(cls, modality: str) -> set[str]:
"""Return supported input formats for a modality.

Format values should be lowercase bare extensions without a leading dot,
such as ``"wav"``, or MIME types, such as ``"audio/wav"``.
"""

return set(getattr(cls, f"{modality}_formats", set()))

def __init__(self, name="", config_root=_config):
self._load_config(config_root)
if "description" not in dir(self):
Expand Down
29 changes: 25 additions & 4 deletions garak/generators/openai.py
Original file line number Diff line number Diff line change
Expand Up @@ -127,8 +127,14 @@
"o1-preview-2024-09-12": 32768,
}

audio_formats = ["wav", "mp3"]
audio_pattern = re.compile("|".join(audio_formats))
audio_mime_subtype_formats = {
"mp3": "mp3",
"mpeg": "mp3",
"wav": "wav",
"x-wav": "wav",
}
# the formats we can send are the mime-map's target values
audio_formats = set(audio_mime_subtype_formats.values())


class OpenAICompatible(Generator):
Expand All @@ -139,6 +145,7 @@ class OpenAICompatible(Generator):
active = True
supports_multiple_generations = False
generator_family_name = "OpenAICompatible" # Placeholder override when extending
audio_formats = audio_formats

# template defaults optionally override when extending
DEFAULT_PARAMS = Generator.DEFAULT_PARAMS | {
Expand Down Expand Up @@ -215,7 +222,7 @@ def _conversation_to_list(conversation: Conversation) -> list[dict]:
},
],
}
elif match := audio_pattern.search(
elif audio_format := audio_mime_subtype_formats.get(
turn.content.data_type[0].split("/")[-1]
):
transformed_turn = {
Expand All @@ -226,7 +233,7 @@ def _conversation_to_list(conversation: Conversation) -> list[dict]:
"type": "input_audio",
"input_audio": {
"data": f"{data_b64}",
"format": match.group(0),
"format": audio_format,
},
},
],
Expand Down Expand Up @@ -376,6 +383,20 @@ def _call_model(
return reponse_message_list


class OpenAIAudioCompatible(OpenAICompatible):
"""OpenAI-compatible chat target explicitly known to accept audio input.

Use this class only for endpoints whose advertised API capability includes
audio. The generic :class:`OpenAICompatible` target remains text-only at
the harness boundary so an arbitrary endpoint is not sent unsupported
binary content.
"""

ENV_VAR = OpenAICompatible.ENV_VAR
generator_family_name = "OpenAIAudioCompatible"
modality = {"in": {"text", "audio"}, "out": {"text"}}


class OpenAIGenerator(OpenAICompatible):
"""Generator wrapper for OpenAI text2text models. Expects API key in the OPENAI_API_KEY environment variable"""

Expand Down
45 changes: 44 additions & 1 deletion tests/generators/test_openai_compatible.py
Original file line number Diff line number Diff line change
Expand Up @@ -11,7 +11,7 @@
from collections.abc import Iterable

from garak.attempt import Message, Turn, Conversation
from garak.generators.openai import OpenAICompatible
from garak.generators.openai import OpenAIAudioCompatible, OpenAICompatible
from garak.generators.rest import RestGenerator

# TODO: expand this when we have faster loading, currently to process all generator costs 30s for 3 tests
Expand Down Expand Up @@ -50,6 +50,13 @@ def compatible() -> Iterable[OpenAICompatible]:
if hasattr(module_klass, "ENV_VAR"):
class_instance = build_test_instance(module_klass)
if isinstance(class_instance, OpenAICompatible):
# this test drives a text prompt; skip generators that do
# not accept text input (e.g. audio-only targets)
modality_in = getattr(class_instance, "modality", {}).get(
"in", {"text"}
)
if "text" not in modality_in:
continue
yield f"{namespace}.{klass_name}"


Expand Down Expand Up @@ -141,3 +148,39 @@ def test_openai_multiple_generations():
assert (
oai_klass.supports_multiple_generations == True
), "OpenAI access expected to correctly support multiple generations by default"


def test_openai_compatible_reports_supported_audio_formats():
assert OpenAICompatible.supported_formats("audio") == {
"wav",
"mp3",
}, "reports audio formats through the generator format interface"
assert (
OpenAICompatible.supported_formats("image") == set()
), "reports no image formats by default"


def test_openai_audio_compatible_declares_audio_modality():
assert OpenAICompatible.modality["in"] == {
"text"
}, "generic compatible targets remain text-only at the harness boundary"
assert OpenAIAudioCompatible.modality["in"] == {
"text",
"audio",
}, "explicit audio-compatible targets accept text plus audio"
assert OpenAIAudioCompatible.supported_formats("audio") == {
"wav",
"mp3",
}, "audio-compatible targets inherit supported wire formats"


def test_openai_compatible_normalises_mp3_audio_payload(tmp_path):
audio_path = tmp_path / "prompt.mp3"
audio_path.write_bytes(b"ID3")
prompt = Conversation([Turn("user", Message("listen", data_path=str(audio_path)))])

payload = OpenAICompatible._conversation_to_list(prompt)

assert (
payload[0]["content"][1]["input_audio"]["format"] == "mp3"
), "normalises audio/mpeg MIME subtype to OpenAI's mp3 format"
Loading