mirror of
https://github.com/HKUDS/nanobot.git
synced 2026-08-08 13:28:43 +03:00
Compare commits
24
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
f44ee6cb27 | ||
|
|
5be3df1d6f | ||
|
|
79c23787f6 | ||
|
|
b647aa5f47 | ||
|
|
a9a8bdcef6 | ||
|
|
2d81cc0ae1 | ||
|
|
aed6b6967c | ||
|
|
3874b3acf4 | ||
|
|
01725bab11 | ||
|
|
a786e3d225 | ||
|
|
626f262121 | ||
|
|
7caf492ae2 | ||
|
|
9aa2ab1657 | ||
|
|
971b774282 | ||
|
|
1377759705 | ||
|
|
d56bafa6d0 | ||
|
|
6ec6c9bb83 | ||
|
|
8a2a5eecdd | ||
|
|
08154b4374 | ||
|
|
880097acd5 | ||
|
|
e02615c93d | ||
|
|
e9259e680e | ||
|
|
a5b85a3d6b | ||
|
|
82c323c2d9 |
@@ -17,6 +17,7 @@ Connect nanobot to your favorite chat platform. Want to build your own? See the
|
|||||||
| **Wecom** | Bot ID + Bot Secret |
|
| **Wecom** | Bot ID + Bot Secret |
|
||||||
| **Microsoft Teams** | App ID + App Password + public HTTPS endpoint |
|
| **Microsoft Teams** | App ID + App Password + public HTTPS endpoint |
|
||||||
| **Mochat** | Claw token (auto-setup available) |
|
| **Mochat** | Claw token (auto-setup available) |
|
||||||
|
| **Signal** | signal-cli daemon + phone number |
|
||||||
|
|
||||||
<details>
|
<details>
|
||||||
<summary><b>Telegram</b> (Recommended)</summary>
|
<summary><b>Telegram</b> (Recommended)</summary>
|
||||||
@@ -669,3 +670,69 @@ nanobot gateway
|
|||||||
```
|
```
|
||||||
|
|
||||||
</details>
|
</details>
|
||||||
|
|
||||||
|
<details>
|
||||||
|
<summary><b>Signal</b></summary>
|
||||||
|
|
||||||
|
Uses **signal-cli** daemon in HTTP mode — receive messages via SSE, send via JSON-RPC.
|
||||||
|
|
||||||
|
**1. Install signal-cli**
|
||||||
|
|
||||||
|
Install [signal-cli](https://github.com/AsamK/signal-cli) and register a phone number:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
signal-cli -u +1234567890 register
|
||||||
|
signal-cli -u +1234567890 verify <CODE>
|
||||||
|
```
|
||||||
|
|
||||||
|
Start the daemon:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
signal-cli -a +1234567890 daemon --http localhost:8080
|
||||||
|
```
|
||||||
|
|
||||||
|
**2. Configure**
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"channels": {
|
||||||
|
"signal": {
|
||||||
|
"enabled": true,
|
||||||
|
"phoneNumber": "+1234567890",
|
||||||
|
"daemonHost": "localhost",
|
||||||
|
"daemonPort": 8080,
|
||||||
|
"dm": {
|
||||||
|
"enabled": true,
|
||||||
|
"policy": "open"
|
||||||
|
},
|
||||||
|
"group": {
|
||||||
|
"enabled": true,
|
||||||
|
"policy": "open",
|
||||||
|
"requireMention": true
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
> - `phoneNumber`: Your registered Signal phone number.
|
||||||
|
> - `daemonHost` / `daemonPort`: Where signal-cli daemon is listening (default `localhost:8080`).
|
||||||
|
> - `dm.policy`: `"open"` (anyone can DM) or `"allowlist"` (only listed numbers/UUIDs). When `"allowlist"`, unlisted DM senders receive a pairing code.
|
||||||
|
> - `dm.allowFrom`: List of allowed phone numbers or UUIDs (used when policy is `"allowlist"`).
|
||||||
|
> - `group.policy`: `"open"` (all groups) or `"allowlist"` (only listed group IDs).
|
||||||
|
> - `group.requireMention`: When `true` (default), the bot only responds in groups when @mentioned.
|
||||||
|
> - `group.allowFrom`: List of allowed group IDs (used when group policy is `"allowlist"`).
|
||||||
|
> - `attachmentsDir`: Override the directory where signal-cli stores inbound attachments. Defaults to `~/.local/share/signal-cli/attachments` (the Linux default). Set this if signal-cli runs with a custom `XDG_DATA_HOME` or on macOS/Windows.
|
||||||
|
> - `groupMessageBufferSize`: Number of recent group messages kept for context (default `20`, must be > 0).
|
||||||
|
|
||||||
|
**3. Run**
|
||||||
|
|
||||||
|
```bash
|
||||||
|
nanobot gateway
|
||||||
|
```
|
||||||
|
|
||||||
|
> [!TIP]
|
||||||
|
> The channel automatically reconnects to the signal-cli daemon with exponential backoff if the connection drops.
|
||||||
|
> Markdown in bot replies is automatically converted to Signal text styles (bold, italic, code, etc.).
|
||||||
|
|
||||||
|
</details>
|
||||||
|
|||||||
@@ -48,6 +48,28 @@ AIHubMix example:
|
|||||||
}
|
}
|
||||||
```
|
```
|
||||||
|
|
||||||
|
Gemini example (Imagen 4):
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"gemini": {
|
||||||
|
"apiKey": "${GEMINI_API_KEY}"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"tools": {
|
||||||
|
"imageGeneration": {
|
||||||
|
"enabled": true,
|
||||||
|
"provider": "gemini",
|
||||||
|
"model": "imagen-4.0-generate-001",
|
||||||
|
"defaultAspectRatio": "1:1"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
For Gemini Flash (which supports reference-image edits) see the [Gemini](#gemini) section below.
|
||||||
|
|
||||||
> [!TIP]
|
> [!TIP]
|
||||||
> Prefer environment variables for API keys. nanobot resolves `${VAR_NAME}` values from the environment at startup.
|
> Prefer environment variables for API keys. nanobot resolves `${VAR_NAME}` values from the environment at startup.
|
||||||
|
|
||||||
@@ -69,7 +91,7 @@ The WebUI hides provider storage details from the user. The agent sees the saved
|
|||||||
| Option | Type | Default | Description |
|
| Option | Type | Default | Description |
|
||||||
|--------|------|---------|-------------|
|
|--------|------|---------|-------------|
|
||||||
| `tools.imageGeneration.enabled` | boolean | `false` | Register the `generate_image` tool |
|
| `tools.imageGeneration.enabled` | boolean | `false` | Register the `generate_image` tool |
|
||||||
| `tools.imageGeneration.provider` | string | `"openrouter"` | Image provider name. Currently `openrouter` and `aihubmix` are supported |
|
| `tools.imageGeneration.provider` | string | `"openrouter"` | Image provider name. Supported values: `openrouter`, `aihubmix`, `gemini` |
|
||||||
| `tools.imageGeneration.model` | string | `"openai/gpt-5.4-image-2"` | Provider model name |
|
| `tools.imageGeneration.model` | string | `"openai/gpt-5.4-image-2"` | Provider model name |
|
||||||
| `tools.imageGeneration.defaultAspectRatio` | string | `"1:1"` | Default ratio when the prompt/tool call does not specify one |
|
| `tools.imageGeneration.defaultAspectRatio` | string | `"1:1"` | Default ratio when the prompt/tool call does not specify one |
|
||||||
| `tools.imageGeneration.defaultImageSize` | string | `"1K"` | Default size hint, for example `1K`, `2K`, `4K`, or `1024x1024` |
|
| `tools.imageGeneration.defaultImageSize` | string | `"1K"` | Default size hint, for example `1K`, `2K`, `4K`, or `1024x1024` |
|
||||||
@@ -139,6 +161,36 @@ Configure:
|
|||||||
|
|
||||||
`quality: low` is optional. It can make free image models faster and less likely to time out, but it is not required for correctness.
|
`quality: low` is optional. It can make free image models faster and less likely to time out, but it is not required for correctness.
|
||||||
|
|
||||||
|
### Gemini
|
||||||
|
|
||||||
|
nanobot supports two Gemini image generation model families via Google's Generative Language API:
|
||||||
|
|
||||||
|
| Model | Endpoint | Reference images |
|
||||||
|
|-------|----------|-----------------|
|
||||||
|
| `imagen-4.0-generate-001` | `:predict` | Not supported by this integration |
|
||||||
|
| `gemini-2.5-flash-image` | `:generateContent` | Supported |
|
||||||
|
|
||||||
|
For reference-image edits, use a Gemini Flash image model:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"gemini": {
|
||||||
|
"apiKey": "${GEMINI_API_KEY}"
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"tools": {
|
||||||
|
"imageGeneration": {
|
||||||
|
"enabled": true,
|
||||||
|
"provider": "gemini",
|
||||||
|
"model": "gemini-2.5-flash-image"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
Imagen 4 supports the aspect ratios `1:1`, `9:16`, `16:9`, `3:4`, and `4:3`. Unsupported ratios are ignored and the model uses its default. The `defaultImageSize` setting has no effect on Gemini models; sizing is controlled by `defaultAspectRatio` only. Reference images passed with an Imagen model are ignored (with a warning logged).
|
||||||
|
|
||||||
## Artifacts
|
## Artifacts
|
||||||
|
|
||||||
Generated images are stored under the active nanobot instance's media directory:
|
Generated images are stored under the active nanobot instance's media directory:
|
||||||
@@ -193,7 +245,7 @@ Use the reference image. Keep the same robot and composition, change the palette
|
|||||||
|---------|-------|
|
|---------|-------|
|
||||||
| `generate_image` is not available | Set `tools.imageGeneration.enabled` to `true` and restart the gateway |
|
| `generate_image` is not available | Set `tools.imageGeneration.enabled` to `true` and restart the gateway |
|
||||||
| Missing API key error | Configure `providers.<provider>.apiKey`; if using `${VAR_NAME}`, confirm the environment variable is visible to the gateway process |
|
| Missing API key error | Configure `providers.<provider>.apiKey`; if using `${VAR_NAME}`, confirm the environment variable is visible to the gateway process |
|
||||||
| `unsupported image generation provider` | Use `openrouter` or `aihubmix` |
|
| `unsupported image generation provider` | Use `openrouter`, `aihubmix`, or `gemini` |
|
||||||
| AIHubMix says `Incorrect model ID` | Use `model: "gpt-image-2-free"`; nanobot expands it to the required `openai/gpt-image-2-free` model path internally |
|
| AIHubMix says `Incorrect model ID` | Use `model: "gpt-image-2-free"`; nanobot expands it to the required `openai/gpt-image-2-free` model path internally |
|
||||||
| Generation times out | Try a smaller/default image size, set AIHubMix `extraBody.quality` to `"low"`, or retry later |
|
| Generation times out | Try a smaller/default image size, set AIHubMix `extraBody.quality` to `"low"`, or retry later |
|
||||||
| Reference image rejected | Reference image paths must be inside the workspace or nanobot media directory and must be valid image files |
|
| Reference image rejected | Reference image paths must be inside the workspace or nanobot media directory and must be valid image files |
|
||||||
|
|||||||
@@ -18,7 +18,9 @@ from nanobot.config.paths import get_media_dir
|
|||||||
from nanobot.config.schema import Base
|
from nanobot.config.schema import Base
|
||||||
from nanobot.providers.image_generation import (
|
from nanobot.providers.image_generation import (
|
||||||
AIHubMixImageGenerationClient,
|
AIHubMixImageGenerationClient,
|
||||||
|
GeminiImageGenerationClient,
|
||||||
ImageGenerationError,
|
ImageGenerationError,
|
||||||
|
MiniMaxImageGenerationClient,
|
||||||
OpenRouterImageGenerationClient,
|
OpenRouterImageGenerationClient,
|
||||||
)
|
)
|
||||||
from nanobot.utils.artifacts import (
|
from nanobot.utils.artifacts import (
|
||||||
@@ -117,7 +119,9 @@ class ImageGenerationTool(Tool):
|
|||||||
def _provider_config(self) -> ProviderConfig | None:
|
def _provider_config(self) -> ProviderConfig | None:
|
||||||
return self.provider_configs.get(self.config.provider)
|
return self.provider_configs.get(self.config.provider)
|
||||||
|
|
||||||
def _provider_client(self) -> OpenRouterImageGenerationClient | AIHubMixImageGenerationClient | None:
|
def _provider_client(
|
||||||
|
self,
|
||||||
|
) -> OpenRouterImageGenerationClient | AIHubMixImageGenerationClient | MiniMaxImageGenerationClient | GeminiImageGenerationClient | None:
|
||||||
provider = self._provider_config()
|
provider = self._provider_config()
|
||||||
kwargs = {
|
kwargs = {
|
||||||
"api_key": provider.api_key if provider else None,
|
"api_key": provider.api_key if provider else None,
|
||||||
@@ -129,6 +133,10 @@ class ImageGenerationTool(Tool):
|
|||||||
return OpenRouterImageGenerationClient(**kwargs)
|
return OpenRouterImageGenerationClient(**kwargs)
|
||||||
if self.config.provider == "aihubmix":
|
if self.config.provider == "aihubmix":
|
||||||
return AIHubMixImageGenerationClient(**kwargs)
|
return AIHubMixImageGenerationClient(**kwargs)
|
||||||
|
if self.config.provider == "minimax":
|
||||||
|
return MiniMaxImageGenerationClient(**kwargs)
|
||||||
|
if self.config.provider == "gemini":
|
||||||
|
return GeminiImageGenerationClient(**kwargs)
|
||||||
return None
|
return None
|
||||||
|
|
||||||
def _missing_api_key_error(self) -> str:
|
def _missing_api_key_error(self) -> str:
|
||||||
@@ -137,6 +145,10 @@ class ImageGenerationTool(Tool):
|
|||||||
return "Error: OpenRouter API key is not configured. Set providers.openrouter.apiKey."
|
return "Error: OpenRouter API key is not configured. Set providers.openrouter.apiKey."
|
||||||
if provider == "aihubmix":
|
if provider == "aihubmix":
|
||||||
return "Error: AIHubMix API key is not configured. Set providers.aihubmix.apiKey."
|
return "Error: AIHubMix API key is not configured. Set providers.aihubmix.apiKey."
|
||||||
|
if provider == "minimax":
|
||||||
|
return "Error: MiniMax API key is not configured. Set providers.minimax.apiKey."
|
||||||
|
if provider == "gemini":
|
||||||
|
return "Error: Gemini API key is not configured. Set providers.gemini.apiKey."
|
||||||
return f"Error: {provider} API key is not configured."
|
return f"Error: {provider} API key is not configured."
|
||||||
|
|
||||||
def _resolve_reference_image(self, value: str) -> str:
|
def _resolve_reference_image(self, value: str) -> str:
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -642,6 +642,8 @@ def serve(
|
|||||||
image_generation_provider_configs={
|
image_generation_provider_configs={
|
||||||
"openrouter": runtime_config.providers.openrouter,
|
"openrouter": runtime_config.providers.openrouter,
|
||||||
"aihubmix": runtime_config.providers.aihubmix,
|
"aihubmix": runtime_config.providers.aihubmix,
|
||||||
|
"minimax": runtime_config.providers.minimax,
|
||||||
|
"gemini": runtime_config.providers.gemini,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
except ValueError as exc:
|
except ValueError as exc:
|
||||||
@@ -755,6 +757,8 @@ def _run_gateway(
|
|||||||
image_generation_provider_configs={
|
image_generation_provider_configs={
|
||||||
"openrouter": config.providers.openrouter,
|
"openrouter": config.providers.openrouter,
|
||||||
"aihubmix": config.providers.aihubmix,
|
"aihubmix": config.providers.aihubmix,
|
||||||
|
"minimax": config.providers.minimax,
|
||||||
|
"gemini": config.providers.gemini,
|
||||||
},
|
},
|
||||||
provider_snapshot_loader=load_provider_snapshot,
|
provider_snapshot_loader=load_provider_snapshot,
|
||||||
runtime_model_publisher=lambda model, preset: publish_runtime_model_update(
|
runtime_model_publisher=lambda model, preset: publish_runtime_model_update(
|
||||||
|
|||||||
@@ -66,6 +66,8 @@ class Nanobot:
|
|||||||
image_generation_provider_configs={
|
image_generation_provider_configs={
|
||||||
"openrouter": config.providers.openrouter,
|
"openrouter": config.providers.openrouter,
|
||||||
"aihubmix": config.providers.aihubmix,
|
"aihubmix": config.providers.aihubmix,
|
||||||
|
"minimax": config.providers.minimax,
|
||||||
|
"gemini": config.providers.gemini,
|
||||||
},
|
},
|
||||||
)
|
)
|
||||||
return cls(loop)
|
return cls(loop)
|
||||||
|
|||||||
@@ -8,6 +8,7 @@ from pathlib import Path
|
|||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
from loguru import logger
|
||||||
|
|
||||||
from nanobot.providers.registry import find_by_name
|
from nanobot.providers.registry import find_by_name
|
||||||
from nanobot.utils.helpers import detect_image_mime
|
from nanobot.utils.helpers import detect_image_mime
|
||||||
@@ -26,6 +27,8 @@ _AIHUBMIX_ASPECT_RATIO_SIZES = {
|
|||||||
"4:3": "1536x1024",
|
"4:3": "1536x1024",
|
||||||
"16:9": "1536x1024",
|
"16:9": "1536x1024",
|
||||||
}
|
}
|
||||||
|
_GEMINI_DEFAULT_TIMEOUT_S = 120.0
|
||||||
|
_GEMINI_IMAGEN_ASPECT_RATIOS = {"1:1", "9:16", "16:9", "3:4", "4:3"}
|
||||||
|
|
||||||
|
|
||||||
class ImageGenerationError(RuntimeError):
|
class ImageGenerationError(RuntimeError):
|
||||||
@@ -50,17 +53,28 @@ def _provider_base_url(provider: str, api_base: str | None, fallback: str) -> st
|
|||||||
return fallback
|
return fallback
|
||||||
|
|
||||||
|
|
||||||
def image_path_to_data_url(path: str | Path) -> str:
|
def _read_image_b64(path: str | Path) -> tuple[str, str]:
|
||||||
"""Convert a local image path to an image data URL."""
|
"""Return ``(mime, base64)`` for the image at ``path``."""
|
||||||
p = Path(path).expanduser()
|
p = Path(path).expanduser()
|
||||||
raw = p.read_bytes()
|
raw = p.read_bytes()
|
||||||
mime = detect_image_mime(raw)
|
mime = detect_image_mime(raw)
|
||||||
if mime is None:
|
if mime is None:
|
||||||
raise ImageGenerationError(f"unsupported reference image: {p}")
|
raise ImageGenerationError(f"unsupported reference image: {p}")
|
||||||
encoded = base64.b64encode(raw).decode("ascii")
|
return mime, base64.b64encode(raw).decode("ascii")
|
||||||
|
|
||||||
|
|
||||||
|
def image_path_to_data_url(path: str | Path) -> str:
|
||||||
|
"""Convert a local image path to an image data URL."""
|
||||||
|
mime, encoded = _read_image_b64(path)
|
||||||
return f"data:{mime};base64,{encoded}"
|
return f"data:{mime};base64,{encoded}"
|
||||||
|
|
||||||
|
|
||||||
|
def image_path_to_inline_data(path: str | Path) -> dict[str, str]:
|
||||||
|
"""Convert a local image path to a Gemini ``inlineData`` payload dict."""
|
||||||
|
mime, encoded = _read_image_b64(path)
|
||||||
|
return {"mimeType": mime, "data": encoded}
|
||||||
|
|
||||||
|
|
||||||
def _b64_png_data_url(value: str) -> str:
|
def _b64_png_data_url(value: str) -> str:
|
||||||
return f"data:image/png;base64,{value}"
|
return f"data:image/png;base64,{value}"
|
||||||
|
|
||||||
@@ -341,6 +355,203 @@ class AIHubMixImageGenerationClient:
|
|||||||
return GeneratedImageResponse(images=images, content="", raw=payload)
|
return GeneratedImageResponse(images=images, content="", raw=payload)
|
||||||
|
|
||||||
|
|
||||||
|
def _http_error_detail(response: httpx.Response) -> str:
|
||||||
|
"""Extract a readable error message from an HTTP error response."""
|
||||||
|
try:
|
||||||
|
data = response.json()
|
||||||
|
if isinstance(data, dict):
|
||||||
|
err = data.get("error")
|
||||||
|
if isinstance(err, dict):
|
||||||
|
return err.get("message") or str(err)
|
||||||
|
if err:
|
||||||
|
return str(err)
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
return response.text[:500] or "<empty response body>"
|
||||||
|
|
||||||
|
|
||||||
|
class GeminiImageGenerationClient:
|
||||||
|
"""Async client for Gemini/Imagen image generation via the Generative Language API."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
api_key: str | None,
|
||||||
|
api_base: str | None = None,
|
||||||
|
extra_headers: dict[str, str] | None = None,
|
||||||
|
extra_body: dict[str, Any] | None = None,
|
||||||
|
timeout: float = _GEMINI_DEFAULT_TIMEOUT_S,
|
||||||
|
client: httpx.AsyncClient | None = None,
|
||||||
|
) -> None:
|
||||||
|
self.api_key = api_key
|
||||||
|
# The Gemini provider's registry default_api_base is the OpenAI-compat
|
||||||
|
# shim (.../v1beta/openai/), which has no image endpoints. Image
|
||||||
|
# generation needs the native Generative Language API base, so we don't
|
||||||
|
# use _provider_base_url() here.
|
||||||
|
self.api_base = (
|
||||||
|
api_base or "https://generativelanguage.googleapis.com/v1beta"
|
||||||
|
).rstrip("/")
|
||||||
|
self.extra_headers = extra_headers or {}
|
||||||
|
self.extra_body = extra_body or {}
|
||||||
|
self.timeout = timeout
|
||||||
|
self._client = client
|
||||||
|
|
||||||
|
async def generate(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
prompt: str,
|
||||||
|
model: str,
|
||||||
|
reference_images: list[str] | None = None,
|
||||||
|
aspect_ratio: str | None = None,
|
||||||
|
image_size: str | None = None,
|
||||||
|
) -> GeneratedImageResponse:
|
||||||
|
if not self.api_key:
|
||||||
|
raise ImageGenerationError(
|
||||||
|
"Gemini API key is not configured. Set providers.gemini.apiKey."
|
||||||
|
)
|
||||||
|
if "imagen" in model.lower():
|
||||||
|
if reference_images:
|
||||||
|
logger.warning(
|
||||||
|
"Imagen models do not support reference images; "
|
||||||
|
"ignoring {} reference image(s) for {}",
|
||||||
|
len(reference_images),
|
||||||
|
model,
|
||||||
|
)
|
||||||
|
return await self._generate_imagen(
|
||||||
|
prompt=prompt, model=model, aspect_ratio=aspect_ratio
|
||||||
|
)
|
||||||
|
return await self._generate_gemini_flash(
|
||||||
|
prompt=prompt, model=model, reference_images=reference_images or []
|
||||||
|
)
|
||||||
|
|
||||||
|
async def _generate_imagen(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
prompt: str,
|
||||||
|
model: str,
|
||||||
|
aspect_ratio: str | None,
|
||||||
|
) -> GeneratedImageResponse:
|
||||||
|
parameters: dict[str, Any] = {"sampleCount": 1}
|
||||||
|
if aspect_ratio in _GEMINI_IMAGEN_ASPECT_RATIOS:
|
||||||
|
parameters["aspectRatio"] = aspect_ratio
|
||||||
|
body: dict[str, Any] = {
|
||||||
|
"instances": [{"prompt": prompt}],
|
||||||
|
"parameters": parameters,
|
||||||
|
}
|
||||||
|
body.update(self.extra_body)
|
||||||
|
|
||||||
|
url = f"{self.api_base}/models/{model}:predict"
|
||||||
|
headers = {
|
||||||
|
"x-goog-api-key": self.api_key or "",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
**self.extra_headers,
|
||||||
|
}
|
||||||
|
|
||||||
|
if self._client is not None:
|
||||||
|
response = await self._client.post(url, headers=headers, json=body)
|
||||||
|
else:
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||||
|
response = await client.post(url, headers=headers, json=body)
|
||||||
|
|
||||||
|
try:
|
||||||
|
response.raise_for_status()
|
||||||
|
except httpx.HTTPStatusError as exc:
|
||||||
|
detail = _http_error_detail(response)
|
||||||
|
logger.error("Gemini Imagen generation failed (HTTP {}): {}", response.status_code, detail)
|
||||||
|
raise ImageGenerationError(
|
||||||
|
f"Gemini Imagen generation failed (HTTP {response.status_code}): {detail}"
|
||||||
|
) from exc
|
||||||
|
|
||||||
|
data = response.json()
|
||||||
|
images: list[str] = []
|
||||||
|
for prediction in data.get("predictions") or []:
|
||||||
|
if not isinstance(prediction, dict):
|
||||||
|
continue
|
||||||
|
b64 = prediction.get("bytesBase64Encoded")
|
||||||
|
mime = prediction.get("mimeType", "image/png")
|
||||||
|
if isinstance(b64, str) and b64:
|
||||||
|
images.append(f"data:{mime};base64,{b64}")
|
||||||
|
|
||||||
|
if not images:
|
||||||
|
provider_error = data.get("error") if isinstance(data, dict) else None
|
||||||
|
if provider_error:
|
||||||
|
raise ImageGenerationError(f"Gemini Imagen returned no images: {provider_error}")
|
||||||
|
raise ImageGenerationError("Gemini Imagen returned no images for this request")
|
||||||
|
|
||||||
|
return GeneratedImageResponse(images=images, content="", raw=data)
|
||||||
|
|
||||||
|
async def _generate_gemini_flash(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
prompt: str,
|
||||||
|
model: str,
|
||||||
|
reference_images: list[str],
|
||||||
|
) -> GeneratedImageResponse:
|
||||||
|
parts: list[dict[str, Any]] = [
|
||||||
|
{"inlineData": image_path_to_inline_data(path)} for path in reference_images
|
||||||
|
]
|
||||||
|
parts.append({"text": prompt})
|
||||||
|
|
||||||
|
body: dict[str, Any] = {
|
||||||
|
"contents": [{"role": "user", "parts": parts}],
|
||||||
|
"generationConfig": {"responseModalities": ["TEXT", "IMAGE"]},
|
||||||
|
}
|
||||||
|
body.update(self.extra_body)
|
||||||
|
|
||||||
|
url = f"{self.api_base}/models/{model}:generateContent"
|
||||||
|
headers = {
|
||||||
|
"x-goog-api-key": self.api_key or "",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
**self.extra_headers,
|
||||||
|
}
|
||||||
|
|
||||||
|
if self._client is not None:
|
||||||
|
response = await self._client.post(url, headers=headers, json=body)
|
||||||
|
else:
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||||
|
response = await client.post(url, headers=headers, json=body)
|
||||||
|
|
||||||
|
try:
|
||||||
|
response.raise_for_status()
|
||||||
|
except httpx.HTTPStatusError as exc:
|
||||||
|
detail = _http_error_detail(response)
|
||||||
|
logger.error("Gemini image generation failed (HTTP {}): {}", response.status_code, detail)
|
||||||
|
raise ImageGenerationError(
|
||||||
|
f"Gemini image generation failed (HTTP {response.status_code}): {detail}"
|
||||||
|
) from exc
|
||||||
|
|
||||||
|
data = response.json()
|
||||||
|
images: list[str] = []
|
||||||
|
text_parts: list[str] = []
|
||||||
|
for candidate in data.get("candidates") or []:
|
||||||
|
if not isinstance(candidate, dict):
|
||||||
|
continue
|
||||||
|
content = candidate.get("content") or {}
|
||||||
|
for part in content.get("parts") or []:
|
||||||
|
if not isinstance(part, dict):
|
||||||
|
continue
|
||||||
|
if "text" in part:
|
||||||
|
text_parts.append(part["text"])
|
||||||
|
inline = part.get("inlineData")
|
||||||
|
if isinstance(inline, dict):
|
||||||
|
mime = inline.get("mimeType", "image/png")
|
||||||
|
b64 = inline.get("data", "")
|
||||||
|
if b64:
|
||||||
|
images.append(f"data:{mime};base64,{b64}")
|
||||||
|
|
||||||
|
if not images:
|
||||||
|
provider_error = data.get("error") if isinstance(data, dict) else None
|
||||||
|
if provider_error:
|
||||||
|
raise ImageGenerationError(f"Gemini returned no images: {provider_error}")
|
||||||
|
raise ImageGenerationError("Gemini returned no images for this request")
|
||||||
|
|
||||||
|
return GeneratedImageResponse(
|
||||||
|
images=images,
|
||||||
|
content="\n".join(t for t in text_parts if t).strip(),
|
||||||
|
raw=data,
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
async def _aihubmix_images_from_payload(
|
async def _aihubmix_images_from_payload(
|
||||||
client: httpx.AsyncClient,
|
client: httpx.AsyncClient,
|
||||||
payload: dict[str, Any],
|
payload: dict[str, Any],
|
||||||
@@ -393,3 +604,136 @@ async def _aihubmix_images_from_payload(
|
|||||||
for candidate in candidates:
|
for candidate in candidates:
|
||||||
await collect(candidate)
|
await collect(candidate)
|
||||||
return images
|
return images
|
||||||
|
|
||||||
|
|
||||||
|
_MINIMAX_TIMEOUT_S = 300.0
|
||||||
|
|
||||||
|
_MINIMAX_ASPECT_RATIO_SIZES = {
|
||||||
|
"1:1": "1:1",
|
||||||
|
"16:9": "16:9",
|
||||||
|
"4:3": "4:3",
|
||||||
|
"3:2": "3:2",
|
||||||
|
"2:3": "2:3",
|
||||||
|
"3:4": "3:4",
|
||||||
|
"9:16": "9:16",
|
||||||
|
"21:9": "21:9",
|
||||||
|
}
|
||||||
|
|
||||||
|
|
||||||
|
class MiniMaxImageGenerationClient:
|
||||||
|
"""Async client for MiniMax image generation API."""
|
||||||
|
|
||||||
|
def __init__(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
api_key: str | None,
|
||||||
|
api_base: str | None = None,
|
||||||
|
extra_headers: dict[str, str] | None = None,
|
||||||
|
extra_body: dict[str, Any] | None = None,
|
||||||
|
timeout: float = _MINIMAX_TIMEOUT_S,
|
||||||
|
client: httpx.AsyncClient | None = None,
|
||||||
|
) -> None:
|
||||||
|
self.api_key = api_key
|
||||||
|
self.api_base = _provider_base_url(
|
||||||
|
"minimax",
|
||||||
|
api_base,
|
||||||
|
"https://api.minimaxi.com/v1",
|
||||||
|
)
|
||||||
|
self.extra_headers = extra_headers or {}
|
||||||
|
self.extra_body = extra_body or {}
|
||||||
|
self.timeout = timeout
|
||||||
|
self._client = client
|
||||||
|
|
||||||
|
def _resolve_aspect_ratio(self, aspect_ratio: str | None) -> str:
|
||||||
|
if aspect_ratio and aspect_ratio in _MINIMAX_ASPECT_RATIO_SIZES:
|
||||||
|
return _MINIMAX_ASPECT_RATIO_SIZES[aspect_ratio]
|
||||||
|
return "1:1"
|
||||||
|
|
||||||
|
async def generate(
|
||||||
|
self,
|
||||||
|
*,
|
||||||
|
prompt: str,
|
||||||
|
model: str,
|
||||||
|
reference_images: list[str] | None = None,
|
||||||
|
aspect_ratio: str | None = None,
|
||||||
|
image_size: str | None = None,
|
||||||
|
) -> GeneratedImageResponse:
|
||||||
|
if not self.api_key:
|
||||||
|
raise ImageGenerationError(
|
||||||
|
"MiniMax API key is not configured. Set providers.minimax.apiKey."
|
||||||
|
)
|
||||||
|
|
||||||
|
headers = {
|
||||||
|
"Authorization": f"Bearer {self.api_key}",
|
||||||
|
"Content-Type": "application/json",
|
||||||
|
**self.extra_headers,
|
||||||
|
}
|
||||||
|
|
||||||
|
body: dict[str, Any] = {
|
||||||
|
"model": model,
|
||||||
|
"prompt": prompt,
|
||||||
|
"response_format": "base64",
|
||||||
|
}
|
||||||
|
|
||||||
|
resolved_ratio = self._resolve_aspect_ratio(aspect_ratio)
|
||||||
|
body["aspect_ratio"] = resolved_ratio
|
||||||
|
|
||||||
|
refs = list(reference_images or [])
|
||||||
|
if refs:
|
||||||
|
image_refs = [image_path_to_data_url(path) for path in refs]
|
||||||
|
body["subject_reference"] = [
|
||||||
|
{"type": "character", "image_file": ref} for ref in image_refs
|
||||||
|
]
|
||||||
|
|
||||||
|
body.update(self.extra_body)
|
||||||
|
|
||||||
|
if self._client is not None:
|
||||||
|
return await self._generate_with_client(self._client, body, headers)
|
||||||
|
async with httpx.AsyncClient(timeout=self.timeout) as client:
|
||||||
|
return await self._generate_with_client(client, body, headers)
|
||||||
|
|
||||||
|
async def _generate_with_client(
|
||||||
|
self,
|
||||||
|
client: httpx.AsyncClient,
|
||||||
|
body: dict[str, Any],
|
||||||
|
headers: dict[str, str],
|
||||||
|
) -> GeneratedImageResponse:
|
||||||
|
url = f"{self.api_base}/image_generation"
|
||||||
|
try:
|
||||||
|
response = await client.post(url, headers=headers, json=body)
|
||||||
|
except httpx.TimeoutException as exc:
|
||||||
|
raise ImageGenerationError("MiniMax image generation timed out") from exc
|
||||||
|
except httpx.RequestError as exc:
|
||||||
|
raise ImageGenerationError(f"MiniMax image generation request failed: {exc}") from exc
|
||||||
|
|
||||||
|
try:
|
||||||
|
response.raise_for_status()
|
||||||
|
except httpx.HTTPStatusError as exc:
|
||||||
|
detail = response.text[:500]
|
||||||
|
raise ImageGenerationError(f"MiniMax image generation failed: {detail}") from exc
|
||||||
|
|
||||||
|
payload = response.json()
|
||||||
|
images = _minimax_images_from_payload(payload)
|
||||||
|
|
||||||
|
if not images:
|
||||||
|
provider_error = payload.get("error") if isinstance(payload, dict) else None
|
||||||
|
if provider_error:
|
||||||
|
raise ImageGenerationError(f"MiniMax returned no images: {provider_error}")
|
||||||
|
raise ImageGenerationError("MiniMax returned no images for this request")
|
||||||
|
|
||||||
|
return GeneratedImageResponse(images=images, content="", raw=payload)
|
||||||
|
|
||||||
|
|
||||||
|
def _minimax_images_from_payload(payload: dict[str, Any]) -> list[str]:
|
||||||
|
"""Extract base64 images from MiniMax API response.
|
||||||
|
|
||||||
|
MiniMax returns images in ``data.image_base64`` (list of base64 strings).
|
||||||
|
"""
|
||||||
|
images: list[str] = []
|
||||||
|
data = payload.get("data")
|
||||||
|
if not isinstance(data, dict):
|
||||||
|
return images
|
||||||
|
for b64 in data.get("image_base64") or []:
|
||||||
|
if isinstance(b64, str) and b64:
|
||||||
|
images.append(_b64_png_data_url(b64))
|
||||||
|
return images
|
||||||
|
|||||||
@@ -88,6 +88,27 @@ AIHubMix `gpt-image-2-free` uses AIHubMix's unified predictions endpoint interna
|
|||||||
|
|
||||||
`providers.aihubmix.extraBody` can be used for provider-specific options. For example, `"extraBody": {"quality": "low"}` is optional but can make `gpt-image-2-free` faster and less likely to time out.
|
`providers.aihubmix.extraBody` can be used for provider-specific options. For example, `"extraBody": {"quality": "low"}` is optional but can make `gpt-image-2-free` faster and less likely to time out.
|
||||||
|
|
||||||
|
For Gemini, the image tool supports two model families. Imagen 4 (`imagen-4.0-generate-001`) supports text-to-image only. Gemini Flash (`gemini-2.5-flash-image`) also supports reference-image edits. Configuration:
|
||||||
|
|
||||||
|
```json
|
||||||
|
{
|
||||||
|
"providers": {
|
||||||
|
"gemini": {
|
||||||
|
"apiKey": "AIza..."
|
||||||
|
}
|
||||||
|
},
|
||||||
|
"tools": {
|
||||||
|
"imageGeneration": {
|
||||||
|
"enabled": true,
|
||||||
|
"provider": "gemini",
|
||||||
|
"model": "imagen-4.0-generate-001"
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
```
|
||||||
|
|
||||||
|
For Gemini models, `defaultImageSize` has no effect; use `defaultAspectRatio` instead. Imagen 4 supports `1:1`, `9:16`, `16:9`, `3:4`, and `4:3`.
|
||||||
|
|
||||||
## Examples
|
## Examples
|
||||||
|
|
||||||
Generate a new image:
|
Generate a new image:
|
||||||
|
|||||||
File diff suppressed because it is too large
Load Diff
@@ -0,0 +1,525 @@
|
|||||||
|
"""Unit tests for the Signal markdown → plain text + textStyle converter."""
|
||||||
|
|
||||||
|
from nanobot.channels.signal import _markdown_to_signal, _partition_styles
|
||||||
|
from nanobot.utils.helpers import split_message
|
||||||
|
|
||||||
|
|
||||||
|
def _utf16_len(s: str) -> int:
|
||||||
|
return len(s.encode("utf-16-le")) // 2
|
||||||
|
|
||||||
|
|
||||||
|
def styles_for(plain: str, text_styles: list[str]) -> dict[str, list[str]]:
|
||||||
|
"""Return a dict mapping each styled substring to its style list."""
|
||||||
|
result: dict[str, list[str]] = {}
|
||||||
|
for entry in text_styles:
|
||||||
|
start_s, length_s, style = entry.split(":", 2)
|
||||||
|
start, length = int(start_s), int(length_s)
|
||||||
|
span = plain[start : start + length]
|
||||||
|
result.setdefault(span, []).append(style)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
def utf16_styles_for(plain: str, text_styles: list[str]) -> dict[str, list[str]]:
|
||||||
|
"""Like styles_for, but slices `plain` using UTF-16 offsets (Signal's units)."""
|
||||||
|
encoded = plain.encode("utf-16-le")
|
||||||
|
result: dict[str, list[str]] = {}
|
||||||
|
for entry in text_styles:
|
||||||
|
start_s, length_s, style = entry.split(":", 2)
|
||||||
|
start, length = int(start_s), int(length_s)
|
||||||
|
span = encoded[start * 2 : (start + length) * 2].decode("utf-16-le")
|
||||||
|
result.setdefault(span, []).append(style)
|
||||||
|
return result
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Basic cases
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_empty():
|
||||||
|
plain, styles = _markdown_to_signal("")
|
||||||
|
assert plain == ""
|
||||||
|
assert styles == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_plain_text():
|
||||||
|
plain, styles = _markdown_to_signal("hello world")
|
||||||
|
assert plain == "hello world"
|
||||||
|
assert styles == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_bold_stars():
|
||||||
|
plain, styles = _markdown_to_signal("say **hello** now")
|
||||||
|
assert plain == "say hello now"
|
||||||
|
assert styles_for(plain, styles) == {"hello": ["BOLD"]}
|
||||||
|
|
||||||
|
|
||||||
|
def test_bold_underscores():
|
||||||
|
plain, styles = _markdown_to_signal("say __hello__ now")
|
||||||
|
assert plain == "say hello now"
|
||||||
|
assert styles_for(plain, styles) == {"hello": ["BOLD"]}
|
||||||
|
|
||||||
|
|
||||||
|
def test_italic_star():
|
||||||
|
plain, styles = _markdown_to_signal("say *hello* now")
|
||||||
|
assert plain == "say hello now"
|
||||||
|
assert styles_for(plain, styles) == {"hello": ["ITALIC"]}
|
||||||
|
|
||||||
|
|
||||||
|
def test_italic_underscore():
|
||||||
|
plain, styles = _markdown_to_signal("say _hello_ now")
|
||||||
|
assert plain == "say hello now"
|
||||||
|
assert styles_for(plain, styles) == {"hello": ["ITALIC"]}
|
||||||
|
|
||||||
|
|
||||||
|
def test_strikethrough():
|
||||||
|
plain, styles = _markdown_to_signal("say ~~hello~~ now")
|
||||||
|
assert plain == "say hello now"
|
||||||
|
assert styles_for(plain, styles) == {"hello": ["STRIKETHROUGH"]}
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Code
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_inline_code():
|
||||||
|
plain, styles = _markdown_to_signal("run `ls -la` here")
|
||||||
|
assert plain == "run ls -la here"
|
||||||
|
assert styles_for(plain, styles) == {"ls -la": ["MONOSPACE"]}
|
||||||
|
|
||||||
|
|
||||||
|
def test_code_block():
|
||||||
|
plain, styles = _markdown_to_signal("```\nprint('hi')\n```")
|
||||||
|
assert "print('hi')" in plain
|
||||||
|
assert styles_for(plain, styles).get("print('hi')\n") == ["MONOSPACE"] or "MONOSPACE" in str(
|
||||||
|
styles_for(plain, styles)
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
|
def test_code_block_with_lang():
|
||||||
|
plain, styles = _markdown_to_signal("```python\ncode\n```")
|
||||||
|
assert "code" in plain
|
||||||
|
assert any("MONOSPACE" in s for s in styles)
|
||||||
|
|
||||||
|
|
||||||
|
def test_code_block_not_processed_further():
|
||||||
|
"""Markdown inside a code block must not be styled."""
|
||||||
|
plain, styles = _markdown_to_signal("```\n**not bold**\n```")
|
||||||
|
assert "**not bold**" in plain
|
||||||
|
# Only MONOSPACE should be applied, no BOLD
|
||||||
|
for entry in styles:
|
||||||
|
assert "BOLD" not in entry
|
||||||
|
|
||||||
|
|
||||||
|
def test_inline_code_not_processed_further():
|
||||||
|
"""Markdown inside inline code must not be styled."""
|
||||||
|
plain, styles = _markdown_to_signal("use `**raw**` please")
|
||||||
|
assert "**raw**" in plain
|
||||||
|
for entry in styles:
|
||||||
|
assert "BOLD" not in entry
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Headers
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_header_becomes_bold():
|
||||||
|
plain, styles = _markdown_to_signal("# My Title")
|
||||||
|
assert plain == "My Title"
|
||||||
|
assert styles_for(plain, styles) == {"My Title": ["BOLD"]}
|
||||||
|
|
||||||
|
|
||||||
|
def test_h2_becomes_bold():
|
||||||
|
plain, styles = _markdown_to_signal("## Sub-section")
|
||||||
|
assert plain == "Sub-section"
|
||||||
|
assert styles_for(plain, styles) == {"Sub-section": ["BOLD"]}
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Blockquotes
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_blockquote_strips_marker():
|
||||||
|
plain, styles = _markdown_to_signal("> some quote")
|
||||||
|
assert plain == "some quote"
|
||||||
|
assert styles == []
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Lists
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_bullet_dash():
|
||||||
|
plain, styles = _markdown_to_signal("- item one")
|
||||||
|
assert plain == "• item one"
|
||||||
|
|
||||||
|
|
||||||
|
def test_bullet_star():
|
||||||
|
plain, styles = _markdown_to_signal("* item two")
|
||||||
|
assert plain == "• item two"
|
||||||
|
|
||||||
|
|
||||||
|
def test_numbered_list():
|
||||||
|
plain, styles = _markdown_to_signal("1. first\n2. second")
|
||||||
|
assert "1. first" in plain
|
||||||
|
assert "2. second" in plain
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Links
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_link_text_differs_from_url():
|
||||||
|
plain, styles = _markdown_to_signal("[Click here](https://example.com)")
|
||||||
|
assert plain == "Click here (https://example.com)"
|
||||||
|
assert styles == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_link_text_equals_url():
|
||||||
|
plain, styles = _markdown_to_signal("[https://example.com](https://example.com)")
|
||||||
|
assert plain == "https://example.com"
|
||||||
|
assert styles == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_link_text_equals_url_without_scheme():
|
||||||
|
plain, styles = _markdown_to_signal("[example.com](https://example.com)")
|
||||||
|
assert plain == "https://example.com"
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Mixed / nesting
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_bold_and_italic_adjacent():
|
||||||
|
plain, styles = _markdown_to_signal("**bold** and *italic*")
|
||||||
|
assert plain == "bold and italic"
|
||||||
|
sd = styles_for(plain, styles)
|
||||||
|
assert sd.get("bold") == ["BOLD"]
|
||||||
|
assert sd.get("italic") == ["ITALIC"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_header_with_inline_code():
|
||||||
|
"""Header becomes BOLD; code inside becomes MONOSPACE (not double-BOLD)."""
|
||||||
|
plain, styles = _markdown_to_signal("# Use `grep`")
|
||||||
|
assert plain == "Use grep"
|
||||||
|
sd = styles_for(plain, styles)
|
||||||
|
assert "BOLD" in sd.get("Use ", []) or "BOLD" in str(styles)
|
||||||
|
assert "MONOSPACE" in sd.get("grep", [])
|
||||||
|
|
||||||
|
|
||||||
|
def test_multiline_mixed():
|
||||||
|
md = "**Title**\n\nSome *italic* text.\n\n- bullet\n- another"
|
||||||
|
plain, styles = _markdown_to_signal(md)
|
||||||
|
assert "Title" in plain
|
||||||
|
assert "italic" in plain
|
||||||
|
assert "• bullet" in plain
|
||||||
|
sd = styles_for(plain, styles)
|
||||||
|
assert "BOLD" in sd.get("Title", [])
|
||||||
|
assert "ITALIC" in sd.get("italic", [])
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Table rendering
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_table_rendered_as_monospace():
|
||||||
|
md = "| A | B |\n| - | - |\n| 1 | 2 |"
|
||||||
|
plain, styles = _markdown_to_signal(md)
|
||||||
|
assert "A" in plain and "B" in plain
|
||||||
|
assert any("MONOSPACE" in s for s in styles)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Style range format
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_style_range_format():
|
||||||
|
"""Each style entry must be 'start:length:STYLE'."""
|
||||||
|
_, styles = _markdown_to_signal("**bold** text")
|
||||||
|
for entry in styles:
|
||||||
|
parts = entry.split(":")
|
||||||
|
assert len(parts) == 3
|
||||||
|
assert parts[0].isdigit()
|
||||||
|
assert parts[1].isdigit()
|
||||||
|
assert parts[2] in {"BOLD", "ITALIC", "STRIKETHROUGH", "MONOSPACE", "SPOILER"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_style_ranges_are_within_bounds():
|
||||||
|
text = "hello **world** end"
|
||||||
|
plain, styles = _markdown_to_signal(text)
|
||||||
|
for entry in styles:
|
||||||
|
start_s, length_s, _ = entry.split(":", 2)
|
||||||
|
start, length = int(start_s), int(length_s)
|
||||||
|
assert start >= 0
|
||||||
|
assert start + length <= len(plain)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Non-BMP / UTF-16 offsets
|
||||||
|
#
|
||||||
|
# Signal's BodyRange (and signal-cli's textStyle) interprets start/length in
|
||||||
|
# UTF-16 code units. Python's len() counts code points, so characters outside
|
||||||
|
# the BMP (emojis, supplementary CJK) shift offsets by +1 per occurrence.
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def assert_within_utf16_bounds(plain: str, styles: list[str]) -> None:
|
||||||
|
limit = _utf16_len(plain)
|
||||||
|
for entry in styles:
|
||||||
|
start_s, length_s, _ = entry.split(":", 2)
|
||||||
|
start, length = int(start_s), int(length_s)
|
||||||
|
assert start >= 0
|
||||||
|
assert start + length <= limit, f"range {entry} exceeds utf-16 length {limit} of {plain!r}"
|
||||||
|
|
||||||
|
|
||||||
|
def test_bold_with_emoji_inside():
|
||||||
|
plain, styles = _markdown_to_signal("**hi 🎉 bye**")
|
||||||
|
assert plain == "hi 🎉 bye"
|
||||||
|
assert utf16_styles_for(plain, styles) == {"hi 🎉 bye": ["BOLD"]}
|
||||||
|
assert_within_utf16_bounds(plain, styles)
|
||||||
|
|
||||||
|
|
||||||
|
def test_italic_with_trailing_emoji():
|
||||||
|
plain, styles = _markdown_to_signal("*bye 🎉*")
|
||||||
|
assert plain == "bye 🎉"
|
||||||
|
assert utf16_styles_for(plain, styles) == {"bye 🎉": ["ITALIC"]}
|
||||||
|
assert_within_utf16_bounds(plain, styles)
|
||||||
|
|
||||||
|
|
||||||
|
def test_bold_after_emoji_prefix():
|
||||||
|
plain, styles = _markdown_to_signal("🎉 **bold**")
|
||||||
|
assert plain == "🎉 bold"
|
||||||
|
assert utf16_styles_for(plain, styles) == {"bold": ["BOLD"]}
|
||||||
|
assert_within_utf16_bounds(plain, styles)
|
||||||
|
|
||||||
|
|
||||||
|
def test_bold_after_and_inside_emoji():
|
||||||
|
plain, styles = _markdown_to_signal("🎉 **a 🎊 b**")
|
||||||
|
assert plain == "🎉 a 🎊 b"
|
||||||
|
assert utf16_styles_for(plain, styles) == {"a 🎊 b": ["BOLD"]}
|
||||||
|
assert_within_utf16_bounds(plain, styles)
|
||||||
|
|
||||||
|
|
||||||
|
def test_supplementary_cjk_in_bold():
|
||||||
|
"""Non-BMP CJK (U+20BB7) proves the bug is UTF-16, not emoji-specific."""
|
||||||
|
plain, styles = _markdown_to_signal("**𠮷野家**")
|
||||||
|
assert plain == "𠮷野家"
|
||||||
|
assert utf16_styles_for(plain, styles) == {"𠮷野家": ["BOLD"]}
|
||||||
|
assert_within_utf16_bounds(plain, styles)
|
||||||
|
|
||||||
|
|
||||||
|
def test_zwj_emoji_in_bold():
|
||||||
|
"""ZWJ family sequence = multiple surrogate pairs + BMP ZWJs."""
|
||||||
|
plain, styles = _markdown_to_signal("**hi 👨👩👧 bye**")
|
||||||
|
assert plain == "hi 👨👩👧 bye"
|
||||||
|
assert utf16_styles_for(plain, styles) == {"hi 👨👩👧 bye": ["BOLD"]}
|
||||||
|
assert_within_utf16_bounds(plain, styles)
|
||||||
|
|
||||||
|
|
||||||
|
def test_ascii_offsets_unchanged():
|
||||||
|
"""ASCII-only path must produce the same offsets as before the UTF-16 fix."""
|
||||||
|
plain, styles = _markdown_to_signal("**bold** plain *it*")
|
||||||
|
assert plain == "bold plain it"
|
||||||
|
assert sorted(styles) == sorted(["0:4:BOLD", "11:2:ITALIC"])
|
||||||
|
|
||||||
|
|
||||||
|
def test_reported_daily_brief_pattern():
|
||||||
|
"""Regression for the reported bug: a single non-BMP emoji shifts every
|
||||||
|
subsequent styled span left by 1 UTF-16 unit, lopping off the last letter.
|
||||||
|
"""
|
||||||
|
md = (
|
||||||
|
"**Weather**\n"
|
||||||
|
"- Conditions: 🌩️ Thunderstorms\n\n"
|
||||||
|
"**News**\n"
|
||||||
|
"*World*\n"
|
||||||
|
"*Local*\n\n"
|
||||||
|
"**Quote of the Day**"
|
||||||
|
)
|
||||||
|
plain, styles = _markdown_to_signal(md)
|
||||||
|
sd = utf16_styles_for(plain, styles)
|
||||||
|
assert sd.get("Weather") == ["BOLD"]
|
||||||
|
assert sd.get("News") == ["BOLD"]
|
||||||
|
assert sd.get("World") == ["ITALIC"]
|
||||||
|
assert sd.get("Local") == ["ITALIC"]
|
||||||
|
assert sd.get("Quote of the Day") == ["BOLD"]
|
||||||
|
assert_within_utf16_bounds(plain, styles)
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Chunk redistribution
|
||||||
|
#
|
||||||
|
# split_message can break a long Signal payload into multiple chunks. The
|
||||||
|
# style ranges from _markdown_to_signal are anchored to the full text, so
|
||||||
|
# they must be redistributed per-chunk with rebased offsets — otherwise
|
||||||
|
# styles for chunks 1..N are silently lost.
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def _resolve_chunk_styles(text: str, max_len: int) -> tuple[list[str], list[list[str]]]:
|
||||||
|
"""Helper: full markdown → signal pipeline, including chunking."""
|
||||||
|
plain, styles = _markdown_to_signal(text)
|
||||||
|
chunks = split_message(plain, max_len) if plain else [""]
|
||||||
|
return chunks, _partition_styles(plain, chunks, styles)
|
||||||
|
|
||||||
|
|
||||||
|
def test_partition_styles_single_chunk_passthrough():
|
||||||
|
plain, styles = _markdown_to_signal("**bold** plain *it*")
|
||||||
|
parts = _partition_styles(plain, [plain], styles)
|
||||||
|
assert parts == [styles]
|
||||||
|
|
||||||
|
|
||||||
|
def test_partition_styles_no_styles():
|
||||||
|
plain = "hello world"
|
||||||
|
assert _partition_styles(plain, [plain], []) == [[]]
|
||||||
|
assert _partition_styles(plain, ["hello", "world"], []) == [[], []]
|
||||||
|
|
||||||
|
|
||||||
|
def test_partition_styles_drops_styles_outside_chunks():
|
||||||
|
"""Whitespace trimmed by split_message must not carry a style range."""
|
||||||
|
plain = "a b"
|
||||||
|
# Fake a style spanning the trimmed whitespace only.
|
||||||
|
chunks = ["a", "b"]
|
||||||
|
parts = _partition_styles(plain, chunks, ["1:3:BOLD"])
|
||||||
|
assert parts == [[], []]
|
||||||
|
|
||||||
|
|
||||||
|
def test_partition_styles_long_message_preserves_chunk_one_styles():
|
||||||
|
"""A bold span deep in the message must follow the message into chunk 1."""
|
||||||
|
# Two ~30-char paragraphs separated by a blank line, then **tail**.
|
||||||
|
line_a = "alpha " * 5 # 30 chars, ends with space
|
||||||
|
line_b = "beta " * 5
|
||||||
|
md = f"{line_a.strip()}\n\n{line_b.strip()}\n\n**tail**"
|
||||||
|
plain, styles = _markdown_to_signal(md)
|
||||||
|
# Force a split between the paragraphs.
|
||||||
|
max_len = len(line_a.strip()) + 2 # fits paragraph A + the "\n\n"
|
||||||
|
chunks = split_message(plain, max_len)
|
||||||
|
assert len(chunks) >= 2, "test setup must produce a split"
|
||||||
|
parts = _partition_styles(plain, chunks, styles)
|
||||||
|
# The bold "tail" should land in the last chunk, with chunk-relative offset.
|
||||||
|
final_chunk = chunks[-1]
|
||||||
|
final_styles = parts[-1]
|
||||||
|
assert any("BOLD" in s for s in final_styles)
|
||||||
|
for entry in final_styles:
|
||||||
|
s, ln, _ = entry.split(":", 2)
|
||||||
|
start, length = int(s), int(ln)
|
||||||
|
slice_ = final_chunk.encode("utf-16-le")[start * 2 : (start + length) * 2].decode(
|
||||||
|
"utf-16-le"
|
||||||
|
)
|
||||||
|
assert slice_ == "tail"
|
||||||
|
|
||||||
|
|
||||||
|
def test_partition_styles_chunk_zero_styles_unchanged():
|
||||||
|
"""Styles entirely in chunk 0 keep their original offsets."""
|
||||||
|
md = "**head** middle and **tail**"
|
||||||
|
plain, styles = _markdown_to_signal(md)
|
||||||
|
# Split so chunk 0 contains "head" and part of the rest, chunk 1 contains "tail".
|
||||||
|
chunks = split_message(plain, 12)
|
||||||
|
assert len(chunks) >= 2
|
||||||
|
parts = _partition_styles(plain, chunks, styles)
|
||||||
|
# "head" lives in chunk 0; assert its offset is unchanged (chunk 0 starts at 0).
|
||||||
|
head_entries = [s for s in parts[0] if "BOLD" in s]
|
||||||
|
assert any(s.startswith("0:4:") for s in head_entries)
|
||||||
|
|
||||||
|
|
||||||
|
def test_partition_styles_with_non_bmp_chunk_offset():
|
||||||
|
"""Chunk-start offsets must be expressed in UTF-16 code units."""
|
||||||
|
# Emoji in chunk 0, bold in chunk 1.
|
||||||
|
md = "🎉 alpha beta gamma\n\n**tail**"
|
||||||
|
plain, styles = _markdown_to_signal(md)
|
||||||
|
chunks = split_message(plain, 18)
|
||||||
|
assert len(chunks) >= 2
|
||||||
|
parts = _partition_styles(plain, chunks, styles)
|
||||||
|
final_styles = parts[-1]
|
||||||
|
assert any("BOLD" in s for s in final_styles)
|
||||||
|
final_chunk = chunks[-1]
|
||||||
|
for entry in final_styles:
|
||||||
|
s, ln, _ = entry.split(":", 2)
|
||||||
|
start, length = int(s), int(ln)
|
||||||
|
slice_ = final_chunk.encode("utf-16-le")[start * 2 : (start + length) * 2].decode(
|
||||||
|
"utf-16-le"
|
||||||
|
)
|
||||||
|
assert slice_ == "tail"
|
||||||
|
|
||||||
|
|
||||||
|
def test_partition_styles_range_spanning_chunks_is_split():
|
||||||
|
"""A style range that straddles a chunk boundary gets sliced into both chunks."""
|
||||||
|
# Construct manually: plain = "abc def", style covers "abc def" (whole thing).
|
||||||
|
plain = "abc def"
|
||||||
|
chunks = split_message(plain, 4) # "abc" / "def"
|
||||||
|
assert chunks == ["abc", "def"]
|
||||||
|
parts = _partition_styles(plain, chunks, ["0:7:BOLD"])
|
||||||
|
# Chunk 0 holds 0:3:BOLD, chunk 1 holds 0:3:BOLD (length=3 each, "def" only
|
||||||
|
# since the space was trimmed by lstrip).
|
||||||
|
assert parts[0] == ["0:3:BOLD"]
|
||||||
|
assert parts[1] == ["0:3:BOLD"]
|
||||||
|
|
||||||
|
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
# Adjacency, nesting, and malformed input
|
||||||
|
# ---------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
def test_bold_italic_combo_outer_bold_inner_italic():
|
||||||
|
"""`**_combo_**` carries both BOLD and ITALIC over the same span."""
|
||||||
|
plain, styles = _markdown_to_signal("**_combo_**")
|
||||||
|
assert plain == "combo"
|
||||||
|
sd = styles_for(plain, styles)
|
||||||
|
assert set(sd.get("combo", [])) == {"BOLD", "ITALIC"}
|
||||||
|
|
||||||
|
|
||||||
|
def test_bold_and_italic_adjacent_no_separator():
|
||||||
|
"""`**bold***italic*` produces BOLD on `bold` and ITALIC on `italic`."""
|
||||||
|
plain, styles = _markdown_to_signal("**bold***italic*")
|
||||||
|
assert plain == "bolditalic"
|
||||||
|
sd = styles_for(plain, styles)
|
||||||
|
assert sd.get("bold") == ["BOLD"]
|
||||||
|
assert sd.get("italic") == ["ITALIC"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_unclosed_bold_falls_through_as_plain():
|
||||||
|
"""An unmatched `**` opener round-trips as literal text with no style."""
|
||||||
|
plain, styles = _markdown_to_signal("**bold")
|
||||||
|
assert plain == "**bold"
|
||||||
|
assert styles == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_unclosed_inline_code_falls_through_as_plain():
|
||||||
|
"""An unmatched backtick round-trips as literal text with no style."""
|
||||||
|
plain, styles = _markdown_to_signal("use `grep")
|
||||||
|
assert plain == "use `grep"
|
||||||
|
assert styles == []
|
||||||
|
|
||||||
|
|
||||||
|
def test_inline_code_inside_blockquote():
|
||||||
|
"""Blockquote prefix is stripped; inline code becomes MONOSPACE."""
|
||||||
|
plain, styles = _markdown_to_signal("> use `grep`")
|
||||||
|
assert plain == "use grep"
|
||||||
|
sd = styles_for(plain, styles)
|
||||||
|
assert sd.get("grep") == ["MONOSPACE"]
|
||||||
|
|
||||||
|
|
||||||
|
def test_header_with_inner_bold_produces_contiguous_bold_ranges():
|
||||||
|
"""`# **wrap** me` — header forces BOLD over the whole line; the inner `**`
|
||||||
|
splits the run, yielding two contiguous BOLD ranges that together cover
|
||||||
|
"wrap me". This is intentional — Signal renders adjacent same-style ranges
|
||||||
|
as a single visual span.
|
||||||
|
"""
|
||||||
|
plain, styles = _markdown_to_signal("# **wrap** me")
|
||||||
|
assert plain == "wrap me"
|
||||||
|
# Both ranges are BOLD; collectively they cover the whole "wrap me".
|
||||||
|
bold_ranges = [s for s in styles if s.endswith(":BOLD")]
|
||||||
|
assert len(bold_ranges) == 2
|
||||||
|
covered = set()
|
||||||
|
for entry in bold_ranges:
|
||||||
|
start, length, _ = entry.split(":", 2)
|
||||||
|
for i in range(int(start), int(start) + int(length)):
|
||||||
|
covered.add(i)
|
||||||
|
assert covered == set(range(len(plain)))
|
||||||
@@ -8,6 +8,7 @@ import pytest
|
|||||||
|
|
||||||
from nanobot.providers.image_generation import (
|
from nanobot.providers.image_generation import (
|
||||||
AIHubMixImageGenerationClient,
|
AIHubMixImageGenerationClient,
|
||||||
|
GeminiImageGenerationClient,
|
||||||
GeneratedImageResponse,
|
GeneratedImageResponse,
|
||||||
ImageGenerationError,
|
ImageGenerationError,
|
||||||
OpenRouterImageGenerationClient,
|
OpenRouterImageGenerationClient,
|
||||||
@@ -202,3 +203,137 @@ async def test_aihubmix_image_generation_downloads_url_response() -> None:
|
|||||||
|
|
||||||
assert response.images[0].startswith("data:image/png;base64,")
|
assert response.images[0].startswith("data:image/png;base64,")
|
||||||
assert fake.get_calls[0]["url"] == "https://cdn.example/image.png"
|
assert fake.get_calls[0]["url"] == "https://cdn.example/image.png"
|
||||||
|
|
||||||
|
|
||||||
|
RAW_B64 = PNG_DATA_URL.removeprefix("data:image/png;base64,")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_gemini_imagen_payload_and_response() -> None:
|
||||||
|
fake = FakeClient(
|
||||||
|
FakeResponse({"predictions": [{"bytesBase64Encoded": RAW_B64, "mimeType": "image/png"}]})
|
||||||
|
)
|
||||||
|
client = GeminiImageGenerationClient(
|
||||||
|
api_key="AIza-test",
|
||||||
|
api_base="https://generativelanguage.googleapis.com/v1beta",
|
||||||
|
client=fake, # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
response = await client.generate(
|
||||||
|
prompt="a sunset",
|
||||||
|
model="imagen-4.0-generate-001",
|
||||||
|
aspect_ratio="16:9",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.images == [PNG_DATA_URL]
|
||||||
|
assert response.content == ""
|
||||||
|
call = fake.calls[0]
|
||||||
|
assert call["url"].endswith(":predict")
|
||||||
|
assert call["headers"]["x-goog-api-key"] == "AIza-test"
|
||||||
|
assert "params" not in call
|
||||||
|
body = call["json"]
|
||||||
|
assert body["instances"] == [{"prompt": "a sunset"}]
|
||||||
|
assert body["parameters"]["sampleCount"] == 1
|
||||||
|
assert body["parameters"]["aspectRatio"] == "16:9"
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_gemini_imagen_ignores_unsupported_aspect_ratio() -> None:
|
||||||
|
fake = FakeClient(
|
||||||
|
FakeResponse({"predictions": [{"bytesBase64Encoded": RAW_B64, "mimeType": "image/png"}]})
|
||||||
|
)
|
||||||
|
client = GeminiImageGenerationClient(api_key="AIza-test", client=fake) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
await client.generate(prompt="a sunset", model="imagen-4.0-generate-001", aspect_ratio="2:3")
|
||||||
|
|
||||||
|
body = fake.calls[0]["json"]
|
||||||
|
assert "aspectRatio" not in body["parameters"]
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_gemini_flash_payload_and_response() -> None:
|
||||||
|
fake = FakeClient(
|
||||||
|
FakeResponse(
|
||||||
|
{
|
||||||
|
"candidates": [
|
||||||
|
{
|
||||||
|
"content": {
|
||||||
|
"parts": [
|
||||||
|
{"text": "here is your image"},
|
||||||
|
{"inlineData": {"mimeType": "image/png", "data": RAW_B64}},
|
||||||
|
]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
)
|
||||||
|
)
|
||||||
|
client = GeminiImageGenerationClient(
|
||||||
|
api_key="AIza-test",
|
||||||
|
api_base="https://generativelanguage.googleapis.com/v1beta",
|
||||||
|
client=fake, # type: ignore[arg-type]
|
||||||
|
)
|
||||||
|
|
||||||
|
response = await client.generate(
|
||||||
|
prompt="draw a cat",
|
||||||
|
model="gemini-2.0-flash-preview-image-generation",
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.images == [PNG_DATA_URL]
|
||||||
|
assert response.content == "here is your image"
|
||||||
|
call = fake.calls[0]
|
||||||
|
assert call["url"].endswith(":generateContent")
|
||||||
|
assert call["headers"]["x-goog-api-key"] == "AIza-test"
|
||||||
|
assert "params" not in call
|
||||||
|
body = call["json"]
|
||||||
|
assert body["generationConfig"]["responseModalities"] == ["TEXT", "IMAGE"]
|
||||||
|
assert body["contents"][0]["parts"][-1] == {"text": "draw a cat"}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_gemini_flash_reference_images(tmp_path: Path) -> None:
|
||||||
|
ref = tmp_path / "ref.png"
|
||||||
|
ref.write_bytes(PNG_BYTES)
|
||||||
|
fake = FakeClient(
|
||||||
|
FakeResponse(
|
||||||
|
{
|
||||||
|
"candidates": [
|
||||||
|
{
|
||||||
|
"content": {
|
||||||
|
"parts": [{"inlineData": {"mimeType": "image/png", "data": RAW_B64}}]
|
||||||
|
}
|
||||||
|
}
|
||||||
|
]
|
||||||
|
}
|
||||||
|
)
|
||||||
|
)
|
||||||
|
client = GeminiImageGenerationClient(api_key="AIza-test", client=fake) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
response = await client.generate(
|
||||||
|
prompt="edit this",
|
||||||
|
model="gemini-2.0-flash-preview-image-generation",
|
||||||
|
reference_images=[str(ref)],
|
||||||
|
)
|
||||||
|
|
||||||
|
assert response.images == [PNG_DATA_URL]
|
||||||
|
parts = fake.calls[0]["json"]["contents"][0]["parts"]
|
||||||
|
assert parts[0]["inlineData"]["mimeType"] == "image/png"
|
||||||
|
assert parts[0]["inlineData"]["data"].startswith("iVBOR")
|
||||||
|
assert parts[1] == {"text": "edit this"}
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_gemini_requires_api_key() -> None:
|
||||||
|
client = GeminiImageGenerationClient(api_key=None)
|
||||||
|
|
||||||
|
with pytest.raises(ImageGenerationError, match="API key"):
|
||||||
|
await client.generate(prompt="draw", model="imagen-4.0-generate-001")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_gemini_no_images_raises() -> None:
|
||||||
|
fake = FakeClient(FakeResponse({"candidates": [{"content": {"parts": [{"text": "sorry"}]}}]}))
|
||||||
|
client = GeminiImageGenerationClient(api_key="AIza-test", client=fake) # type: ignore[arg-type]
|
||||||
|
|
||||||
|
with pytest.raises(ImageGenerationError, match="returned no images"):
|
||||||
|
await client.generate(prompt="draw", model="gemini-2.0-flash-preview-image-generation")
|
||||||
|
|||||||
Reference in New Issue
Block a user