Files
vibe-bot/vibe_bot/tests/test_services.py
T
2026-08-19 13:16:43 -04:00

954 lines
32 KiB
Python

"""Service-layer tests: chat, image, speech, and conversation services.
These exercise the LLM-backed logic directly (constructing each service with
mock dependencies) rather than going through the thin Discord command wrappers.
"""
from __future__ import annotations
import asyncio
import base64
from io import BytesIO
from typing import Any
from unittest.mock import AsyncMock, MagicMock, patch
import pytest
import requests
from vibe_bot.config import TTS_VOICE
from vibe_bot.services.chat_service import ChatService
from vibe_bot.services.conversation_service import (
MAX_TOPIC_LENGTH,
ConversationService,
flip_counter,
)
from vibe_bot.services.image_service import (
MAX_IMAGE_DOWNLOAD_BYTES,
MAX_IMAGE_PROMPT_LENGTH,
ImageService,
_allowed_image_url,
_download_image_bytes,
select_image_layout,
verify_image_prompt,
)
from vibe_bot.services.speech_service import (
MAX_SPEAK_LENGTH,
SpeechService,
parse_voice_flag,
)
@pytest.fixture
def mock_ctx() -> MagicMock:
"""A mock Discord command context."""
ctx = MagicMock()
ctx.author.name = "testuser"
ctx.author.id = "12345"
ctx.author.global_name = "Test User"
ctx.author.nick = "tester"
ctx.author.top_role.name = "@everyone"
ctx.author.activities = []
ctx.author.joined_at = None
ctx.author.created_at = None
ctx.channel.id = "channel-1"
ctx.guild.id = "guild-1"
ctx.message.id = "msg-1"
ctx.message.attachments = []
ctx.bot.user = MagicMock()
ctx.bot.user.name = "test-bot"
ctx.bot.user.id = "bot-123"
ctx.send = AsyncMock()
return ctx
def _file_factory() -> MagicMock:
"""A File factory that records (buffer, filename) as a tuple."""
factory = MagicMock()
def make_file(buf: BytesIO, name: str) -> tuple[str, BytesIO, str]:
return ("FILE", buf, name)
factory.side_effect = make_file
return factory
def _sent_texts(ctx: MagicMock) -> list[str]:
"""All positional text messages sent through ctx.send."""
return [c.args[0] for c in ctx.send.call_args_list if c.args]
def _registry() -> MagicMock:
"""A mock ToolRegistry."""
reg = MagicMock()
reg.to_openai_tools.return_value = []
reg.execute.return_value = "tool result"
return reg
def _fake_bot(name: str) -> tuple[str, str, str, str]:
"""A stand-in custom bot tuple for manager.get_custom_bot."""
return (name, "a personality", "user-123", "2024-01-01")
# ---------------------------------------------------------------------------
# ChatService
# ---------------------------------------------------------------------------
def test_chat_success(mock_ctx: MagicMock) -> None:
"""A normal turn persists the exchange and sends the reply."""
db = MagicMock()
db.get_conversation_context.return_value = []
svc = ChatService(db, _registry())
with patch(
"vibe_bot.llm_client.chat_completion_with_tools",
new=AsyncMock(return_value="This is a bot response"),
):
asyncio.run(
svc.handle(
mock_ctx,
bot_name="alfred",
message="hello",
system_prompt="you are a butler",
response_prefix="alfred response",
)
)
db.add_message.assert_called()
assert mock_ctx.send.call_count >= 2
def test_chat_turn_embedding_budget(
mock_ctx: MagicMock,
temp_db_path: str,
) -> None:
"""One chat turn embeds exactly twice: the RAG query and the user row.
The assistant row is persisted with embed=False, so it costs no
embedding call and stores no embedding row.
"""
import sqlite3
from vibe_bot.database import ChatDatabase
db = ChatDatabase(db_path=temp_db_path)
svc = ChatService(db, _registry())
with (
patch(
"vibe_bot.llm_client.embedding",
return_value=[0.25] * 32,
) as mock_embedding,
patch(
"vibe_bot.llm_client.chat_completion_with_tools",
new=AsyncMock(return_value="This is a bot response"),
),
):
asyncio.run(
svc.handle(
mock_ctx,
bot_name="alfred",
message="hello",
system_prompt="you are a butler",
response_prefix="alfred response",
)
)
assert mock_embedding.call_count == 2
conn = sqlite3.connect(temp_db_path)
embedding_rows = conn.execute("SELECT COUNT(*) FROM message_embeddings").fetchone()
conn.close()
assert embedding_rows[0] == 1
def test_chat_error(mock_ctx: MagicMock) -> None:
"""An LLM error surfaces a friendly message."""
db = MagicMock()
db.get_conversation_context.return_value = []
svc = ChatService(db, _registry())
with patch(
"vibe_bot.llm_client.chat_completion_with_tools",
new=AsyncMock(side_effect=Exception("API error")),
):
asyncio.run(
svc.handle(
mock_ctx,
bot_name="alfred",
message="hello",
system_prompt="you are a butler",
response_prefix="alfred response",
)
)
call_args = mock_ctx.send.call_args[0][0]
assert "error occurred" in call_args.lower()
db.add_message.assert_not_called()
def test_chat_long_response_chunked(mock_ctx: MagicMock) -> None:
"""Long responses are split into multiple sends."""
db = MagicMock()
db.get_conversation_context.return_value = []
svc = ChatService(db, _registry())
with patch(
"vibe_bot.llm_client.chat_completion_with_tools",
new=AsyncMock(return_value="x" * 2500),
):
asyncio.run(
svc.handle(
mock_ctx,
bot_name="alfred",
message="hello",
system_prompt="you are a butler",
response_prefix="alfred response",
)
)
assert mock_ctx.send.call_count >= 3
def test_chat_includes_user_info(mock_ctx: MagicMock) -> None:
"""The system prompt sent to the LLM includes the requester's info."""
db = MagicMock()
db.get_conversation_context.return_value = []
svc = ChatService(db, _registry())
mock_llm = AsyncMock(return_value="resp")
with patch("vibe_bot.llm_client.chat_completion_with_tools", new=mock_llm):
asyncio.run(
svc.handle(
mock_ctx,
bot_name="alfred",
message="hello",
system_prompt="you are a butler",
response_prefix="alfred response",
)
)
system_prompt = mock_llm.call_args.kwargs["system_prompt"]
assert "testuser" in system_prompt
def test_chat_with_context_and_tools(mock_ctx: MagicMock) -> None:
"""Prior RAG context is prepended and tool calls reach the registry."""
db = MagicMock()
db.get_conversation_context.return_value = [
{"role": "user", "content": "old question"},
{"role": "assistant", "content": "old answer"},
]
registry = _registry()
captured: dict[str, Any] = {}
async def fake_llm(**kwargs: Any) -> str:
captured.update(kwargs)
kwargs["tool_executor"]("get_channel_members", {})
await kwargs["tool_call_notifier"]("get_channel_members", {})
return "resp"
with patch("vibe_bot.llm_client.chat_completion_with_tools", new=fake_llm):
asyncio.run(
ChatService(db, registry).handle(
mock_ctx,
bot_name="alfred",
message="hello",
system_prompt="you are a butler",
response_prefix="alfred response",
)
)
prompts = captured["prompts"]
assert prompts[0] == {"role": "user", "content": "old question"}
assert prompts[-1] == {"role": "user", "content": "hello"}
registry.execute.assert_called_once_with(
"get_channel_members", {}, channel=mock_ctx.channel
)
assert any("looking at the channel members" in t for t in _sent_texts(mock_ctx))
# ---------------------------------------------------------------------------
# SpeechService
# ---------------------------------------------------------------------------
def _speech_service(
tts: MagicMock | None,
manager: MagicMock | None = None,
make_file: MagicMock | None = None,
) -> SpeechService:
return SpeechService(
MagicMock(), manager or MagicMock(), tts, make_file or _file_factory()
)
def test_speak_tts_not_initialized(mock_ctx: MagicMock) -> None:
"""No TTS engine means a clear error, no LLM or TTS calls."""
svc = _speech_service(None)
asyncio.run(svc.speak(mock_ctx, message="hello world"))
call_args = mock_ctx.send.call_args[0][0]
assert "TTS engine not initialized" in call_args
def test_speak_empty_message(mock_ctx: MagicMock) -> None:
"""Empty text is rejected before any TTS work."""
svc = _speech_service(MagicMock())
asyncio.run(svc.speak(mock_ctx, message=""))
call_args = mock_ctx.send.call_args[0][0]
assert "Please provide text" in call_args
def test_speak_too_long(mock_ctx: MagicMock) -> None:
"""Oversized text is rejected without calling the TTS engine."""
tts = MagicMock()
svc = _speech_service(tts)
asyncio.run(svc.speak(mock_ctx, message="a" * (MAX_SPEAK_LENGTH + 1)))
tts.generate_audio.assert_not_called()
call_args = mock_ctx.send.call_args[0][0]
assert "Text too long to speak" in call_args
def test_speak_partial_audio_warns(mock_ctx: MagicMock) -> None:
"""Partial audio triggers a warning line."""
tts = MagicMock()
tts.generate_audio.return_value = MagicMock(
audio=MagicMock(), partial=True, failed_chunks=1
)
manager = MagicMock()
manager.list_custom_bots.return_value = []
svc = _speech_service(tts, manager=manager)
asyncio.run(svc.speak(mock_ctx, message="hello world"))
assert any("audio may be incomplete" in t for t in _sent_texts(mock_ctx))
def test_speak_plain_text(mock_ctx: MagicMock) -> None:
"""Plain text is spoken and the audio file is sent."""
tts = MagicMock()
tts.generate_audio.return_value = MagicMock(audio=MagicMock(), partial=False)
manager = MagicMock()
manager.list_custom_bots.return_value = []
svc = _speech_service(tts, manager=manager)
asyncio.run(svc.speak(mock_ctx, message="hello world"))
tts.generate_audio.assert_called_once()
assert mock_ctx.send.call_count >= 2
def test_speak_with_custom_bot(mock_ctx: MagicMock) -> None:
"""A bot prefix routes through the LLM, then speaks the response."""
tts = MagicMock()
tts.generate_audio.return_value = MagicMock(audio=MagicMock(), partial=False)
manager = MagicMock()
manager.list_custom_bots.return_value = [
("alfred", "british butler", "user-123"),
]
manager.get_custom_bot.return_value = (
"alfred",
"british butler",
"user-123",
"2024-01-01",
)
svc = _speech_service(tts, manager=manager)
with patch(
"vibe_bot.llm_client.chat_completion_with_tools",
new=AsyncMock(return_value="The time is 3pm"),
):
asyncio.run(svc.speak(mock_ctx, message="alfred what time is it"))
tts.generate_audio.assert_called_once()
assert any("**alfred**:" in t for t in _sent_texts(mock_ctx))
def test_speak_uses_requested_voice(mock_ctx: MagicMock) -> None:
"""A trailing --voice flag selects that voice for the TTS call."""
tts = MagicMock()
tts.generate_audio.return_value = MagicMock(audio=MagicMock(), partial=False)
manager = MagicMock()
manager.list_custom_bots.return_value = []
svc = _speech_service(tts, manager=manager)
asyncio.run(svc.speak(mock_ctx, message="hello world --voice af_bella"))
assert tts.generate_audio.call_args.kwargs["voice"] == "af_bella"
def test_speak_mid_text_voice_flag_spoken_verbatim(mock_ctx: MagicMock) -> None:
"""A --voice mid-message is preserved as speech and the default voice is used."""
tts = MagicMock()
tts.generate_audio.return_value = MagicMock(audio=MagicMock(), partial=False)
manager = MagicMock()
manager.list_custom_bots.return_value = []
svc = _speech_service(tts, manager=manager)
message = "hello --voice af_bella world"
asyncio.run(svc.speak(mock_ctx, message=message))
assert tts.generate_audio.call_args.args[0] == message
assert tts.generate_audio.call_args.kwargs["voice"] == TTS_VOICE
def test_speak_unknown_voice(mock_ctx: MagicMock) -> None:
"""An unknown voice is rejected before any TTS call."""
tts = MagicMock()
manager = MagicMock()
manager.list_custom_bots.return_value = []
svc = _speech_service(tts, manager=manager)
asyncio.run(svc.speak(mock_ctx, message="hello --voice not_a_real_voice"))
tts.generate_audio.assert_not_called()
call_args = mock_ctx.send.call_args[0][0]
assert "Unknown voice" in call_args
def test_speak_language_lookup_uses_precomputed_dict(
mock_ctx: MagicMock,
) -> None:
"""The speak hot path resolves the language via VOICE_LANGUAGES.get().
The dict is built once at import from VOICES_LIST (covering every
catalog voice) and the per-speak lookup is a single dict get — no
per-call scan of the category list.
"""
from vibe_bot.config import VOICES_LIST
from vibe_bot.services import speech_service
class LanguageLookupSpy:
"""Counts .get() lookups on the voice->language mapping."""
def __init__(self, data: dict[str, str]) -> None:
self.data = data
self.lookups = 0
def get(self, key: str, default: str | None = None) -> str | None:
self.lookups += 1
return self.data.get(key, default)
def __contains__(self, key: object) -> bool:
return key in self.data
counting = LanguageLookupSpy(speech_service.VOICE_LANGUAGES)
assert set(counting.data) == {
voice for category in VOICES_LIST.values() for voice in category["voices"]
}
tts = MagicMock()
tts.generate_audio.return_value = MagicMock(audio=MagicMock(), partial=False)
manager = MagicMock()
manager.list_custom_bots.return_value = []
svc = _speech_service(tts, manager=manager)
with patch.object(speech_service, "VOICE_LANGUAGES", counting):
asyncio.run(svc.speak(mock_ctx, message="hello --voice bf_alice"))
assert counting.lookups == 1
assert tts.generate_audio.call_args.kwargs["lang"] == "en-gb"
# ---------------------------------------------------------------------------
# ConversationService
# ---------------------------------------------------------------------------
def _conversation_service(
manager: MagicMock | None = None,
) -> ConversationService:
return ConversationService(manager or MagicMock())
def test_flip_counter() -> None:
"""flip_counter toggles between 0 and 1."""
assert flip_counter(0) == 1
assert flip_counter(1) == 0
def test_talkforme_topic_too_long(mock_ctx: MagicMock) -> None:
"""Oversized topics are rejected before any LLM call."""
svc = _conversation_service()
asyncio.run(svc.run(mock_ctx, "a", "b", "3", "x" * (MAX_TOPIC_LENGTH + 1)))
call_args = mock_ctx.send.call_args[0][0]
assert "Topic too long" in call_args
def test_talkforme_bot1_not_found(mock_ctx: MagicMock) -> None:
"""A missing first bot is reported and the run stops."""
manager = MagicMock()
manager.get_custom_bot.return_value = None
svc = _conversation_service(manager=manager)
asyncio.run(svc.run(mock_ctx, "ghost", "alfred", "3", "cats"))
call_args = mock_ctx.send.call_args[0][0]
assert "ghost is not a real bot" in call_args
def test_talkforme_invalid_limit(mock_ctx: MagicMock) -> None:
"""A non-integer limit is rejected after both bots are found."""
manager = MagicMock()
manager.get_custom_bot.side_effect = _fake_bot
svc = _conversation_service(manager=manager)
asyncio.run(svc.run(mock_ctx, "a", "b", "abc", "cats"))
call_args = mock_ctx.send.call_args[0][0]
assert "Message limit must be an integer" in call_args
def test_talkforme_first_reply_chunked(mock_ctx: MagicMock) -> None:
"""Long first replies are sent in multiple chunks."""
manager = MagicMock()
manager.get_custom_bot.side_effect = _fake_bot
svc = _conversation_service(manager=manager)
with patch(
"vibe_bot.llm_client.chat_completion_with_history",
new=AsyncMock(return_value="y" * 2500),
):
asyncio.run(svc.run(mock_ctx, "a", "b", "1", "cats"))
assert mock_ctx.send.call_count >= 3
# ---------------------------------------------------------------------------
# ImageService (doodlebob / retcon)
# ---------------------------------------------------------------------------
def _image_service(db: MagicMock | None = None) -> ImageService:
return ImageService(db or MagicMock(), _file_factory())
def test_doodlebob_prompt_too_long(mock_ctx: MagicMock) -> None:
"""Oversized prompts are rejected before any LLM call."""
svc = _image_service()
asyncio.run(svc.generate(mock_ctx, message="a" * (MAX_IMAGE_PROMPT_LENGTH + 1)))
call_args = mock_ctx.send.call_args[0][0]
assert "Prompt too long" in call_args
def test_doodlebob_generate_success(mock_ctx: MagicMock) -> None:
"""A full generate flow ends with an image file and a completion line."""
db = MagicMock()
db.get_image_generation_time_estimate.return_value = None
svc = _image_service(db)
b64 = base64.b64encode(b"fake image").decode()
with (
patch(
"vibe_bot.llm_client.chat_completion_instruct",
new=AsyncMock(side_effect=["square", "a detailed prompt", "PASS"]),
),
patch(
"vibe_bot.llm_client.image_generation",
new=AsyncMock(return_value=b64),
),
):
asyncio.run(svc.generate(mock_ctx, message="a centaur in a field"))
assert db.record_image_generation_time.called
assert any("Strike complete" in t for t in _sent_texts(mock_ctx))
def test_doodlebob_failed_generation(mock_ctx: MagicMock) -> None:
"""An empty image-generation response is reported as a failure."""
db = MagicMock()
db.get_image_generation_time_estimate.return_value = None
svc = _image_service(db)
with (
patch(
"vibe_bot.llm_client.chat_completion_instruct",
new=AsyncMock(side_effect=["square", "a detailed prompt", "PASS"]),
),
patch(
"vibe_bot.llm_client.image_generation",
new=AsyncMock(return_value=""),
),
):
asyncio.run(svc.generate(mock_ctx, message="a centaur"))
assert any("Failed to generate image" in t for t in _sent_texts(mock_ctx))
assert not db.record_image_generation_time.called
def test_doodlebob_reports_estimate(mock_ctx: MagicMock) -> None:
"""A prior-history estimate produces a Drone ETA line."""
db = MagicMock()
db.get_image_generation_time_estimate.return_value = 12.5
svc = _image_service(db)
b64 = base64.b64encode(b"fake image").decode()
with (
patch(
"vibe_bot.llm_client.chat_completion_instruct",
new=AsyncMock(side_effect=["square", "a detailed prompt", "PASS"]),
),
patch(
"vibe_bot.llm_client.image_generation",
new=AsyncMock(return_value=b64),
),
):
asyncio.run(svc.generate(mock_ctx, message="a centaur"))
assert any("Drone ETA" in t for t in _sent_texts(mock_ctx))
def test_doodlebob_no_estimate_without_history(mock_ctx: MagicMock) -> None:
"""No estimate means no Drone ETA line."""
db = MagicMock()
db.get_image_generation_time_estimate.return_value = None
svc = _image_service(db)
b64 = base64.b64encode(b"fake image").decode()
with (
patch(
"vibe_bot.llm_client.chat_completion_instruct",
new=AsyncMock(side_effect=["square", "a detailed prompt", "PASS"]),
),
patch(
"vibe_bot.llm_client.image_generation",
new=AsyncMock(return_value=b64),
),
):
asyncio.run(svc.generate(mock_ctx, message="a centaur"))
assert not any("Drone ETA" in t for t in _sent_texts(mock_ctx))
def test_doodlebob_empty_prompt_stops(mock_ctx: MagicMock) -> None:
"""An empty image-prompt response stops the flow without generating."""
db = MagicMock()
svc = _image_service(db)
with (
patch(
"vibe_bot.llm_client.chat_completion_instruct",
new=AsyncMock(return_value=""),
),
patch("vibe_bot.llm_client.image_generation") as mock_gen,
):
asyncio.run(svc.generate(mock_ctx, message="a centaur"))
mock_gen.assert_not_called()
assert not db.record_image_generation_time.called
def test_doodlebob_decode_failure(mock_ctx: MagicMock) -> None:
"""Invalid base64 from the generation API is reported as a failure."""
db = MagicMock()
db.get_image_generation_time_estimate.return_value = None
svc = _image_service(db)
with (
patch(
"vibe_bot.llm_client.chat_completion_instruct",
new=AsyncMock(side_effect=["square", "a detailed prompt", "PASS"]),
),
patch(
"vibe_bot.llm_client.image_generation",
new=AsyncMock(return_value="abcde!!!"),
),
):
asyncio.run(svc.generate(mock_ctx, message="a centaur"))
assert any(
"Failed to process the generated image" in t for t in _sent_texts(mock_ctx)
)
@pytest.mark.parametrize(
("response", "expected"),
[
("portrait", "portrait"),
("landscape", "landscape"),
("square", "square"),
("PORTRAIT", "portrait"),
("I think landscape", "landscape"),
],
)
def test_select_image_layout_returns_parsed(
response: str,
expected: str,
) -> None:
"""select_image_layout parses the LLM's layout choice."""
with patch(
"vibe_bot.llm_client.chat_completion_instruct",
new=AsyncMock(return_value=response),
):
result = asyncio.run(select_image_layout("a tall tree"))
assert result == expected
def test_select_image_layout_uses_minimal_token_budget() -> None:
"""Layout selection is a one-word answer, so max_tokens is 2."""
mock_llm = AsyncMock(return_value="square")
with patch("vibe_bot.llm_client.chat_completion_instruct", new=mock_llm):
assert asyncio.run(select_image_layout("a tall tree")) == "square"
assert mock_llm.call_args.kwargs["max_tokens"] == 2
def test_doodlebob_latency_within_budget(mock_ctx: MagicMock) -> None:
"""End-to-end doodlebob latency with 50ms simulated per LLM/image call.
Hermetic latency figure: four mocked calls (layout, prompt, verify,
generate) at 50ms each must dominate the wall time; the overhead on
top of the simulated 200ms stays far below the budget.
"""
import time
db = MagicMock()
db.get_image_generation_time_estimate.return_value = None
svc = _image_service(db)
b64 = base64.b64encode(b"fake image").decode()
responses = ["square", "a detailed prompt", "PASS"]
async def slow_instruct(**_kwargs: Any) -> str:
await asyncio.sleep(0.05)
return responses.pop(0)
async def slow_generate(**_kwargs: Any) -> str:
await asyncio.sleep(0.05)
return b64
with (
patch("vibe_bot.llm_client.chat_completion_instruct", new=slow_instruct),
patch("vibe_bot.llm_client.image_generation", new=slow_generate),
):
start = time.monotonic()
asyncio.run(svc.generate(mock_ctx, message="a centaur in a field"))
elapsed = time.monotonic() - start
assert elapsed >= 0.2
assert elapsed < 5.0
assert any("Strike complete" in t for t in _sent_texts(mock_ctx))
def test_verify_image_prompt_pass_keeps_prompt() -> None:
"""A PASS verdict keeps the original prompt unchanged."""
with patch(
"vibe_bot.llm_client.chat_completion_instruct",
new=AsyncMock(return_value="PASS"),
):
result = asyncio.run(verify_image_prompt("a centaur", "a detailed prompt"))
assert result == "a detailed prompt"
def test_verify_image_prompt_correction_replaces() -> None:
"""A non-passing verdict is used as the corrected prompt."""
correction = "a rewritten prompt that is definitely long enough to be a fix"
with patch(
"vibe_bot.llm_client.chat_completion_instruct",
new=AsyncMock(return_value=correction),
):
result = asyncio.run(verify_image_prompt("a centaur", "a detailed prompt"))
assert result == correction
def test_verify_image_prompt_empty_falls_back() -> None:
"""An empty verdict falls back to the original prompt."""
with patch(
"vibe_bot.llm_client.chat_completion_instruct",
new=AsyncMock(return_value=""),
):
result = asyncio.run(verify_image_prompt("a centaur", "a detailed prompt"))
assert result == "a detailed prompt"
def test_retcon_no_attachments(mock_ctx: MagicMock) -> None:
"""retcon with no attachments asks the user to attach an image."""
svc = _image_service()
mock_ctx.message.attachments = []
asyncio.run(svc.edit(mock_ctx, message="make it blue"))
call_args = mock_ctx.send.call_args[0][0]
assert "Please attach an image" in call_args
def test_retcon_rejected_url_not_downloaded(mock_ctx: MagicMock) -> None:
"""A non-Discord attachment URL is refused before any download happens."""
svc = _image_service()
attachment = MagicMock()
attachment.url = "https://evil.example.com/img.png"
mock_ctx.message.attachments = [attachment]
mock_edit = AsyncMock(return_value="")
with (
patch("vibe_bot.llm_client.image_edit", new=mock_edit),
patch("requests.get") as mock_get,
):
asyncio.run(svc.edit(mock_ctx, message="make it blue"))
mock_get.assert_not_called()
mock_edit.assert_not_called()
assert any("Please attach an image" in t for t in _sent_texts(mock_ctx))
def test_retcon_prompt_too_long(mock_ctx: MagicMock) -> None:
"""Oversized retcon prompts are rejected before any download."""
svc = _image_service()
asyncio.run(svc.edit(mock_ctx, message="a" * (MAX_IMAGE_PROMPT_LENGTH + 1)))
call_args = mock_ctx.send.call_args[0][0]
assert "Prompt too long" in call_args
def test_retcon_image_edit_empty(mock_ctx: MagicMock) -> None:
"""An empty edit response is reported as a failure."""
svc = _image_service()
attachment = MagicMock()
attachment.url = "https://cdn.discordapp.com/attachments/1/2/3/img.png"
mock_ctx.message.attachments = [attachment]
with (
patch(
"vibe_bot.services.image_service._download_image_bytes",
return_value=b"fake image bytes",
),
patch(
"vibe_bot.llm_client.image_edit",
new=AsyncMock(return_value=""),
),
):
asyncio.run(svc.edit(mock_ctx, message="make it blue"))
call_args = mock_ctx.send.call_args[0][0]
assert "Failed to edit the image" in call_args
def test_retcon_success(mock_ctx: MagicMock) -> None:
"""A successful edit sends the edited image file."""
svc = _image_service()
attachment = MagicMock()
attachment.url = "https://cdn.discordapp.com/attachments/1/2/3/img.png"
mock_ctx.message.attachments = [attachment]
b64 = base64.b64encode(b"edited").decode()
with (
patch(
"vibe_bot.services.image_service._download_image_bytes",
return_value=b"fake image bytes",
),
patch(
"vibe_bot.llm_client.image_edit",
new=AsyncMock(return_value=b64),
),
):
asyncio.run(svc.edit(mock_ctx, message="make it blue"))
assert any("Rewriting history" in t for t in _sent_texts(mock_ctx))
def test_retcon_edit_decode_failure(mock_ctx: MagicMock) -> None:
"""Invalid base64 from the edit API is reported as a processing failure."""
svc = _image_service()
attachment = MagicMock()
attachment.url = "https://cdn.discordapp.com/attachments/1/2/3/img.png"
mock_ctx.message.attachments = [attachment]
with (
patch(
"vibe_bot.services.image_service._download_image_bytes",
return_value=b"fake image bytes",
),
patch(
"vibe_bot.llm_client.image_edit",
new=AsyncMock(return_value="abcde!!!"),
),
):
asyncio.run(svc.edit(mock_ctx, message="make it blue"))
call_args = mock_ctx.send.call_args[0][0]
assert "Failed to process the edited image" in call_args
def test_allowed_image_url() -> None:
"""Only Discord CDN hosts are allowed for retcon downloads."""
assert _allowed_image_url("https://cdn.discordapp.com/a/b/c.png")
assert _allowed_image_url("https://media.discordapp.net/a/b/c.png")
assert not _allowed_image_url("https://example.com/a/b/c.png")
assert not _allowed_image_url("https://evilcdn.com/a/b/c.png")
assert not _allowed_image_url("http://[::1")
def test_download_image_bytes_success() -> None:
"""An allowed Discord URL is downloaded and its chunks joined."""
response = MagicMock()
response.headers = {}
response.iter_content.return_value = [b"abc", b"", b"def"]
response.raise_for_status.return_value = None
with patch(
"vibe_bot.services.image_service.requests.get",
return_value=response,
) as mock_get:
data = _download_image_bytes("https://cdn.discordapp.com/a/b/c.png")
mock_get.assert_called_once()
assert data == b"abcdef"
def test_download_image_bytes_request_failure() -> None:
"""A failing download returns None instead of raising."""
with patch(
"vibe_bot.services.image_service.requests.get",
side_effect=requests.RequestException("boom"),
):
assert _download_image_bytes("https://cdn.discordapp.com/a/b/c.png") is None
def test_download_image_bytes_content_length_cap() -> None:
"""A declared Content-Length above the cap is refused before streaming."""
response = MagicMock()
response.headers = {"Content-Length": str(MAX_IMAGE_DOWNLOAD_BYTES + 1)}
response.iter_content = MagicMock()
with patch("vibe_bot.services.image_service.requests.get", return_value=response):
assert _download_image_bytes("https://cdn.discordapp.com/a/b/c.png") is None
response.iter_content.assert_not_called()
def test_download_image_bytes_streaming_cap() -> None:
"""Streaming past the size cap is refused even without a Content-Length."""
response = MagicMock()
response.headers = {}
response.iter_content.return_value = [b"a" * (MAX_IMAGE_DOWNLOAD_BYTES + 1)]
with patch("vibe_bot.services.image_service.requests.get", return_value=response):
assert _download_image_bytes("https://cdn.discordapp.com/a/b/c.png") is None
def test_download_image_bytes_bad_content_length() -> None:
"""A non-numeric Content-Length is ignored and streaming proceeds."""
response = MagicMock()
response.headers = {"Content-Length": "not-a-number"}
response.iter_content.return_value = [b"xyz"]
response.raise_for_status.return_value = None
with patch("vibe_bot.services.image_service.requests.get", return_value=response):
assert _download_image_bytes("https://cdn.discordapp.com/a/b/c.png") == b"xyz"
def test_download_image_bytes_stream_error() -> None:
"""A mid-stream request error returns None instead of raising."""
response = MagicMock()
response.headers = {}
response.iter_content.side_effect = requests.RequestException("stream died")
with patch("vibe_bot.services.image_service.requests.get", return_value=response):
assert _download_image_bytes("https://cdn.discordapp.com/a/b/c.png") is None
# ---------------------------------------------------------------------------
# parse_voice_flag
# ---------------------------------------------------------------------------
def test_parse_voice_flag_no_flag() -> None:
"""Without a trailing flag the message is returned unchanged."""
assert parse_voice_flag("hello world") == ("hello world", None)
def test_parse_voice_flag_trailing() -> None:
"""A trailing --voice flag is split off."""
assert parse_voice_flag("hello world --voice af_bella") == (
"hello world",
"af_bella",
)
def test_parse_voice_flag_mid_text_preserved() -> None:
"""A --voice that is not at the end is treated as speech, not a flag."""
assert parse_voice_flag("hello --voice af_bella world") == (
"hello --voice af_bella world",
None,
)
def test_parse_voice_flag_missing_value() -> None:
"""A --voice with no value (or only trailing spaces) is not a flag."""
assert parse_voice_flag("hello --voice") == ("hello --voice", None)
assert parse_voice_flag("hello --voice ") == ("hello --voice ", None)