add time estimation for image gen
Build and Push Container / build-and-push (push) Successful in 1m11s

This commit is contained in:
2026-08-17 12:29:26 -04:00
parent a73a10fdb6
commit b740d00b5d
4 changed files with 275 additions and 0 deletions
+97
View File
@@ -30,6 +30,11 @@ logging.basicConfig(
) )
logger = logging.getLogger(__name__) logger = logging.getLogger(__name__)
# Moving-average window (most recent generations) for the time estimate.
IMAGE_GEN_TIME_WINDOW = 10
# Maximum number of generation times retained in the database.
IMAGE_GEN_TIME_LIMIT = 100
class ChatDatabase: class ChatDatabase:
"""SQLite database with RAG support for storing chat history """SQLite database with RAG support for storing chat history
@@ -118,6 +123,19 @@ class ChatDatabase:
) )
logger.info("idx_user_id index created successfully") logger.info("idx_user_id index created successfully")
# Create image generation timing table
logger.info("Creating image_generation_times table if not exists")
cursor.execute(
"""
CREATE TABLE IF NOT EXISTS image_generation_times (
id INTEGER PRIMARY KEY AUTOINCREMENT,
duration_seconds REAL NOT NULL,
generated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
)
""",
)
logger.info("image_generation_times table initialized successfully")
conn.commit() conn.commit()
logger.info("Database initialization completed successfully") logger.info("Database initialization completed successfully")
conn.close() conn.close()
@@ -497,6 +515,85 @@ class ChatDatabase:
conn.commit() conn.commit()
conn.close() conn.close()
def record_image_generation_time(self, duration_seconds: float) -> bool:
"""Record how long an image generation took.
Args:
duration_seconds: Wall-clock seconds the generation took.
"""
logger.info("Recording image generation time: %.2fs", duration_seconds)
conn = sqlite3.connect(self.db_path)
cursor = conn.cursor()
try:
cursor.execute(
"""
INSERT INTO image_generation_times (duration_seconds)
VALUES (?)
""",
(duration_seconds,),
)
# Cap the table so it doesn't grow unbounded.
cursor.execute(
"""
DELETE FROM image_generation_times
WHERE id NOT IN (
SELECT id FROM image_generation_times
ORDER BY id DESC
LIMIT ?
)
""",
(IMAGE_GEN_TIME_LIMIT,),
)
conn.commit()
except Exception:
logger.exception("Error recording image generation time")
conn.rollback()
return False
else:
return True
finally:
conn.close()
def get_image_generation_time_estimate(self) -> float | None:
"""Get a moving-average estimate of image generation time.
Returns:
The average duration in seconds over the most recent generations,
or None if there is no generation history yet.
"""
logger.debug("Computing moving-average image generation time estimate")
conn = sqlite3.connect(self.db_path)
cursor = conn.cursor()
try:
cursor.execute(
"""
SELECT AVG(duration_seconds)
FROM (
SELECT duration_seconds
FROM image_generation_times
ORDER BY id DESC
LIMIT ?
)
""",
(IMAGE_GEN_TIME_WINDOW,),
)
row = cursor.fetchone()
except Exception:
logger.exception("Error reading image generation times")
return None
finally:
conn.close()
if row is None or row[0] is None:
return None
return float(row[0])
# Global database instance # Global database instance
_chat_db: ChatDatabase | None = None _chat_db: ChatDatabase | None = None
+12
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
import base64 import base64
import logging import logging
import re import re
import time
from io import BytesIO from io import BytesIO
from typing import TYPE_CHECKING from typing import TYPE_CHECKING
@@ -819,8 +820,13 @@ async def doodlebob(ctx: CommandsContext[Bot], *, message: str) -> None:
return return
# Alert the user we're generating the image # Alert the user we're generating the image
db = get_database()
estimated_seconds = db.get_image_generation_time_estimate()
await ctx.send(f"**Doodlebob calling drone strike on {image_prompt[:100]}...**") await ctx.send(f"**Doodlebob calling drone strike on {image_prompt[:100]}...**")
if estimated_seconds is not None:
await ctx.send(f"**Drone ETA: ~{estimated_seconds:.0f} seconds**")
start_time = time.monotonic()
image_b64 = llama_wrapper.image_generation( image_b64 = llama_wrapper.image_generation(
prompt=image_prompt, prompt=image_prompt,
openai_url=IMAGE_GEN_ENDPOINT, openai_url=IMAGE_GEN_ENDPOINT,
@@ -828,16 +834,22 @@ async def doodlebob(ctx: CommandsContext[Bot], *, message: str) -> None:
model=IMAGE_GEN_MODEL, model=IMAGE_GEN_MODEL,
size=LAYOUT_SIZES[layout], size=LAYOUT_SIZES[layout],
) )
elapsed_seconds = time.monotonic() - start_time
if not image_b64: if not image_b64:
logger.warning("Image generation returned empty response.") logger.warning("Image generation returned empty response.")
await ctx.send("Failed to generate image. The server may be busy.") await ctx.send("Failed to generate image. The server may be busy.")
return return
db.record_image_generation_time(elapsed_seconds)
try: try:
edited_image_data = BytesIO(base64.b64decode(image_b64)) edited_image_data = BytesIO(base64.b64decode(image_b64))
send_img = discord.File(edited_image_data, filename="image.png") send_img = discord.File(edited_image_data, filename="image.png")
await ctx.send(file=send_img) await ctx.send(file=send_img)
await ctx.send(
f"**Strike complete. Image generated in {elapsed_seconds:.1f} seconds.**",
)
except Exception: except Exception:
logger.exception("Failed to decode image data") logger.exception("Failed to decode image data")
await ctx.send("Failed to process the generated image.") await ctx.send("Failed to process the generated image.")
+71
View File
@@ -194,6 +194,76 @@ def test_clear_all_messages(
assert len(messages) == 0 assert len(messages) == 0
def test_image_generation_time_estimate_empty(chat_db: ChatDatabase) -> None:
"""Test the estimate is None before any generations are recorded."""
assert chat_db.get_image_generation_time_estimate() is None
def test_record_and_get_image_generation_time_estimate(
chat_db: ChatDatabase,
) -> None:
"""Test recorded generation times produce an average estimate."""
assert chat_db.record_image_generation_time(10.0) is True
assert chat_db.record_image_generation_time(20.0) is True
estimate = chat_db.get_image_generation_time_estimate()
assert estimate == pytest.approx(15.0)
def test_image_generation_estimate_uses_recent_window(
chat_db: ChatDatabase,
) -> None:
"""Test only the most recent window of generations feeds the estimate."""
from vibe_bot.database import IMAGE_GEN_TIME_WINDOW
chat_db.record_image_generation_time(1000.0)
for _ in range(IMAGE_GEN_TIME_WINDOW):
chat_db.record_image_generation_time(10.0)
estimate = chat_db.get_image_generation_time_estimate()
assert estimate == pytest.approx(10.0)
def test_image_generation_times_capped(chat_db: ChatDatabase) -> None:
"""Test the generation time table keeps only the most recent entries."""
import sqlite3
from vibe_bot.database import IMAGE_GEN_TIME_LIMIT
total = IMAGE_GEN_TIME_LIMIT + 5
for i in range(total):
chat_db.record_image_generation_time(float(i))
conn = sqlite3.connect(chat_db.db_path)
cursor = conn.cursor()
cursor.execute("SELECT COUNT(*) FROM image_generation_times")
count = cursor.fetchone()[0]
cursor.execute("SELECT MIN(id) FROM image_generation_times")
min_id = cursor.fetchone()[0]
cursor.execute("SELECT MAX(id) FROM image_generation_times")
max_id = cursor.fetchone()[0]
conn.close()
assert count == IMAGE_GEN_TIME_LIMIT
assert min_id == 6
assert max_id == total
def test_image_generation_estimate_after_capping(
chat_db: ChatDatabase,
) -> None:
"""Test the estimate is computed from the retained recent rows."""
from vibe_bot.database import IMAGE_GEN_TIME_LIMIT
total = IMAGE_GEN_TIME_LIMIT + 5
for i in range(total):
chat_db.record_image_generation_time(float(i))
# Rows 95.0 through 104.0 are the ten most recent retained generations.
estimate = chat_db.get_image_generation_time_estimate()
assert estimate == pytest.approx(99.5)
def test_get_user_history( def test_get_user_history(
chat_db: ChatDatabase, chat_db: ChatDatabase,
mock_embedding: MagicMock, mock_embedding: MagicMock,
@@ -484,3 +554,4 @@ def test_database_init_creates_tables(temp_db_path: str) -> None:
assert "chat_messages" in tables assert "chat_messages" in tables
assert "message_embeddings" in tables assert "message_embeddings" in tables
assert "custom_bots" in tables assert "custom_bots" in tables
assert "image_generation_times" in tables
+95
View File
@@ -1022,6 +1022,7 @@ def test_doodlebob_selects_portrait(
mock_ctx: MagicMock, mock_ctx: MagicMock,
mock_llama_wrapper: MagicMock, mock_llama_wrapper: MagicMock,
mock_base64: MagicMock, mock_base64: MagicMock,
mock_database: MagicMock,
) -> None: ) -> None:
"""Test doodlebob picks portrait and passes the portrait size.""" """Test doodlebob picks portrait and passes the portrait size."""
import asyncio import asyncio
@@ -1033,6 +1034,7 @@ def test_doodlebob_selects_portrait(
"a tall portrait of a lighthouse", # prompt rewrite "a tall portrait of a lighthouse", # prompt rewrite
] ]
mock_llama_wrapper.image_generation.return_value = "aW1hZ2U=" mock_llama_wrapper.image_generation.return_value = "aW1hZ2U="
mock_database.get_image_generation_time_estimate.return_value = None
with patch.object(main_module, "LAYOUT_SIZES", LAYOUT_TEST_SIZES): with patch.object(main_module, "LAYOUT_SIZES", LAYOUT_TEST_SIZES):
asyncio.run(main_module.doodlebob(mock_ctx, message="a tall lighthouse")) asyncio.run(main_module.doodlebob(mock_ctx, message="a tall lighthouse"))
@@ -1055,6 +1057,7 @@ def test_doodlebob_selects_landscape(
mock_ctx: MagicMock, mock_ctx: MagicMock,
mock_llama_wrapper: MagicMock, mock_llama_wrapper: MagicMock,
mock_base64: MagicMock, mock_base64: MagicMock,
mock_database: MagicMock,
) -> None: ) -> None:
"""Test doodlebob picks landscape and passes the landscape size.""" """Test doodlebob picks landscape and passes the landscape size."""
import asyncio import asyncio
@@ -1066,6 +1069,7 @@ def test_doodlebob_selects_landscape(
"a wide panoramic coastline", "a wide panoramic coastline",
] ]
mock_llama_wrapper.image_generation.return_value = "aW1hZ2U=" mock_llama_wrapper.image_generation.return_value = "aW1hZ2U="
mock_database.get_image_generation_time_estimate.return_value = None
with patch.object(main_module, "LAYOUT_SIZES", LAYOUT_TEST_SIZES): with patch.object(main_module, "LAYOUT_SIZES", LAYOUT_TEST_SIZES):
asyncio.run(main_module.doodlebob(mock_ctx, message="wide coastline")) asyncio.run(main_module.doodlebob(mock_ctx, message="wide coastline"))
@@ -1077,6 +1081,7 @@ def test_doodlebob_malformed_layout_defaults_square(
mock_ctx: MagicMock, mock_ctx: MagicMock,
mock_llama_wrapper: MagicMock, mock_llama_wrapper: MagicMock,
mock_base64: MagicMock, mock_base64: MagicMock,
mock_database: MagicMock,
) -> None: ) -> None:
"""Test a malformed layout response falls back to the square size.""" """Test a malformed layout response falls back to the square size."""
import asyncio import asyncio
@@ -1088,6 +1093,7 @@ def test_doodlebob_malformed_layout_defaults_square(
"a balanced composition", "a balanced composition",
] ]
mock_llama_wrapper.image_generation.return_value = "aW1hZ2U=" mock_llama_wrapper.image_generation.return_value = "aW1hZ2U="
mock_database.get_image_generation_time_estimate.return_value = None
with patch.object(main_module, "LAYOUT_SIZES", LAYOUT_TEST_SIZES): with patch.object(main_module, "LAYOUT_SIZES", LAYOUT_TEST_SIZES):
asyncio.run(main_module.doodlebob(mock_ctx, message="a logo")) asyncio.run(main_module.doodlebob(mock_ctx, message="a logo"))
@@ -1101,6 +1107,7 @@ def test_doodlebob_empty_layout_defaults_square(
mock_ctx: MagicMock, mock_ctx: MagicMock,
mock_llama_wrapper: MagicMock, mock_llama_wrapper: MagicMock,
mock_base64: MagicMock, mock_base64: MagicMock,
mock_database: MagicMock,
) -> None: ) -> None:
"""Test an empty layout response (LLM failure) falls back to square.""" """Test an empty layout response (LLM failure) falls back to square."""
import asyncio import asyncio
@@ -1112,8 +1119,96 @@ def test_doodlebob_empty_layout_defaults_square(
"a balanced composition", "a balanced composition",
] ]
mock_llama_wrapper.image_generation.return_value = "aW1hZ2U=" mock_llama_wrapper.image_generation.return_value = "aW1hZ2U="
mock_database.get_image_generation_time_estimate.return_value = None
with patch.object(main_module, "LAYOUT_SIZES", LAYOUT_TEST_SIZES): with patch.object(main_module, "LAYOUT_SIZES", LAYOUT_TEST_SIZES):
asyncio.run(main_module.doodlebob(mock_ctx, message="a logo")) asyncio.run(main_module.doodlebob(mock_ctx, message="a logo"))
assert mock_llama_wrapper.image_generation.call_args.kwargs["size"] == "1024x1024" assert mock_llama_wrapper.image_generation.call_args.kwargs["size"] == "1024x1024"
def test_doodlebob_reports_estimate_and_elapsed(
mock_ctx: MagicMock,
mock_llama_wrapper: MagicMock,
mock_base64: MagicMock,
mock_database: MagicMock,
) -> None:
"""Doodlebob posts an ETA from history and the final elapsed time."""
import asyncio
import re
import vibe_bot.main as main_module
mock_llama_wrapper.chat_completion_instruct.side_effect = [
"square",
"a test scene",
]
mock_llama_wrapper.image_generation.return_value = "aW1hZ2U="
mock_database.get_image_generation_time_estimate.return_value = 12.34
with patch.object(main_module, "LAYOUT_SIZES", LAYOUT_TEST_SIZES):
asyncio.run(main_module.doodlebob(mock_ctx, message="a scene"))
mock_database.record_image_generation_time.assert_called_once()
recorded = mock_database.record_image_generation_time.call_args[0][0]
assert isinstance(recorded, float)
assert recorded >= 0.0
sent = _sent_texts(mock_ctx)
assert any("drone strike" in m for m in sent)
assert any("~12 seconds" in m for m in sent)
assert any(re.search(r"generated in \d+\.\d+ seconds", m) for m in sent)
def test_doodlebob_no_estimate_without_history(
mock_ctx: MagicMock,
mock_llama_wrapper: MagicMock,
mock_base64: MagicMock,
mock_database: MagicMock,
) -> None:
"""No ETA is posted when there is no generation history yet."""
import asyncio
import vibe_bot.main as main_module
mock_llama_wrapper.chat_completion_instruct.side_effect = [
"square",
"a test scene",
]
mock_llama_wrapper.image_generation.return_value = "aW1hZ2U="
mock_database.get_image_generation_time_estimate.return_value = None
with patch.object(main_module, "LAYOUT_SIZES", LAYOUT_TEST_SIZES):
asyncio.run(main_module.doodlebob(mock_ctx, message="a scene"))
sent = _sent_texts(mock_ctx)
assert any("drone strike" in m for m in sent)
assert not any("ETA" in m for m in sent)
mock_database.record_image_generation_time.assert_called_once()
def test_doodlebob_failed_generation_not_recorded(
mock_ctx: MagicMock,
mock_llama_wrapper: MagicMock,
mock_base64: MagicMock,
mock_database: MagicMock,
) -> None:
"""A failed generation records no time and reports no elapsed seconds."""
import asyncio
import vibe_bot.main as main_module
mock_llama_wrapper.chat_completion_instruct.side_effect = [
"square",
"a test scene",
]
mock_llama_wrapper.image_generation.return_value = ""
mock_database.get_image_generation_time_estimate.return_value = 8.0
with patch.object(main_module, "LAYOUT_SIZES", LAYOUT_TEST_SIZES):
asyncio.run(main_module.doodlebob(mock_ctx, message="a scene"))
mock_database.record_image_generation_time.assert_not_called()
sent = _sent_texts(mock_ctx)
assert any("Failed to generate image" in m for m in sent)
assert not any("generated in" in m for m in sent)