Rework AI-Typewriter en application d'arrière-plan multi-profils
- Rollback gestion d'équations : écriture littérale caractère par caractère (suppression de math_format.py et de la normalisation Unicode). - Nouveau profil par défaut 'Mathématiques (LaTeX)' : le modèle émet du LaTeX encadré [EQ]...[/EQ], intercepté par KeyStepper qui déclenche Alt+= en début d'équation et -> en fin. - Config multi-profils dans %APPDATA%/ai-typewriter/config.json (ConfigStore + profils par défaut auto-créés). - Icône de zone de notification (pystray) : Ouvrir les logs (temps réel), Ajouter/Modifier un profil, Gérer l'authentification, Quitter. - Sélecteur de modèles : modèles locaux Ollama + recherche/téléchargement depuis la bibliothèque publique ; fournisseurs tiers (OpenAI, OpenRouter, Gemini, custom). - Clés d'API stockées de façon sécurisée dans le Gestionnaire d'identifiants Windows via keyring (jamais en clair dans config.json). - Tests : 44 tests verts (stepper, profils, clients IA, credentials, catalogue de modèles, moteur).
This commit is contained in:
@@ -12,6 +12,7 @@ dist/
|
||||
|
||||
# Local secrets/config
|
||||
config.json
|
||||
*.key
|
||||
|
||||
# Binary releases
|
||||
*.exe
|
||||
|
||||
@@ -1,132 +1,101 @@
|
||||
# ai-typewriter
|
||||
# AI-Typewriter
|
||||
|
||||
Application Python qui lit la dernière entrée texte du presse-papier avec un raccourci global, l'envoie à un modèle IA, puis remplace chaque pression de touche suivante par le caractère suivant de la réponse.
|
||||
Application Python qui tourne **en arrière-plan** (icône dans la zone de notification de Windows, aucune fenêtre au démarrage), lit le contenu du presse-papier sur un raccourci global, l'envoie à un modèle IA, puis réécrit la réponse **caractère par caractère** à chaque pression de touche physique.
|
||||
|
||||
Un profil spécialisé « Mathématiques (LaTeX) » marque les équations avec `[EQ]...[/EQ]` : l'application les intercepte et déclenche `Alt+=` pour ouvrir une équation (Word/OneNote) et `→` pour en sortir.
|
||||
|
||||
## Fonctionnement
|
||||
|
||||
1. L'utilisateur copie manuellement le texte à envoyer à l'IA.
|
||||
1. L'utilisateur copie le texte à envoyer à l'IA.
|
||||
2. Raccourci global par défaut : `Ctrl+Alt+A`.
|
||||
3. L'application lit directement la dernière entrée du presse-papier, sans simuler `Ctrl+C`.
|
||||
4. Le texte capturé est journalisé puis envoyé à Ollama ou Gemini.
|
||||
5. Le temps de génération de la réponse est journalisé.
|
||||
6. Quand la réponse arrive, le mode dactylographie s'active.
|
||||
7. Chaque touche physique appuyée est interceptée et remplacée par le prochain caractère de la réponse IA.
|
||||
8. Le hook clavier est libéré automatiquement après le dernier caractère.
|
||||
|
||||
L'application n'interprète pas les équations et ne lance pas `Alt+=`. Elle écrit uniquement le texte généré, caractère par caractère. Pour les maths, le modèle peut produire du LaTeX encadré par des marqueurs `[EQ]...[/EQ]`, puis la gestion Word peut être faite ailleurs.
|
||||
3. L'application lit le presse-papier et envoie au modèle du profil actif.
|
||||
4. Une icône reste disponible dans la zone de notification : elle permet d'ouvrir les journaux, d'ajouter/commuter des profils et de gérer les clés d'API.
|
||||
5. Chaque touche physique appuyée ensuite écrit l'élément suivant de la réponse (caractère ou séquence d'équation).
|
||||
6. Le hook clavier est libéré automatiquement à la fin de la réponse.
|
||||
|
||||
## Installation depuis les sources
|
||||
|
||||
```bash
|
||||
python -m venv .venv
|
||||
. .venv/bin/activate
|
||||
. .venv/bin/activate # Windows : .venv\Scripts\activate
|
||||
pip install -r requirements.txt
|
||||
cp config.json.template config.json
|
||||
pip install -e .
|
||||
python main.py
|
||||
```
|
||||
|
||||
Sous Linux, le paquet `keyboard` nécessite souvent les droits root ou l'accès aux périphériques `/dev/input`. Sous Windows, lancez l'exécutable dans une session utilisateur normale.
|
||||
À la première exécution, l'application crée son fichier de configuration :
|
||||
`%APPDATA%\ai-typewriter\config.json` (Linux : `~/.config/ai-typewriter/config.json`).
|
||||
|
||||
## Configuration
|
||||
## Configuration (profils)
|
||||
|
||||
Copiez `config.json.template` vers `config.json` puis adaptez :
|
||||
La configuration contient une liste de **profils** nommés et le profil actif. Chaque profil décrit :
|
||||
|
||||
```json
|
||||
{
|
||||
"provider": "ollama",
|
||||
"model": "llama3.1",
|
||||
"api_key": "",
|
||||
"server_url": "http://localhost:11434",
|
||||
"hotkey": "ctrl+alt+a",
|
||||
"request_timeout_seconds": 300,
|
||||
"math_text_format": "plain"
|
||||
}
|
||||
```
|
||||
| Champ | Description |
|
||||
|---|---|
|
||||
| `name` | Libellé affiché dans les menus |
|
||||
| `provider` | `ollama`, `openai`, `openrouter`, `gemini`, `custom` (OpenAI-compatible) |
|
||||
| `model` | Nom du modèle (choisissable via le sélecteur) |
|
||||
| `server_url` | Base de l'instance (ex. `http://localhost:11434`) |
|
||||
| `credential` | Nom logique de la clé d'API (voir « Authentification ») |
|
||||
| `system_prompt` | Instructions données au modèle |
|
||||
| `equation_enabled` | Active l'interception des marqueurs d'équation |
|
||||
| `eq_start_marker` / `eq_end_marker` | Marqueurs (défaut `[EQ]` / `[/EQ]`) |
|
||||
| `eq_start_key` / `eq_end_key` | Touches déclenchées (défaut `alt+=` / `right`) |
|
||||
|
||||
### Ollama
|
||||
Deux profils sont créés par défaut : **Général** et **Mathématiques (LaTeX)**.
|
||||
|
||||
```json
|
||||
{
|
||||
"provider": "ollama",
|
||||
"model": "llama3.1",
|
||||
"server_url": "http://localhost:11434"
|
||||
}
|
||||
```
|
||||
### Profil Mathématiques (LaTeX)
|
||||
|
||||
### Gemini
|
||||
|
||||
```json
|
||||
{
|
||||
"provider": "gemini",
|
||||
"model": "gemini-1.5-flash",
|
||||
"api_key": "VOTRE_CLE",
|
||||
"server_url": "https://generativelanguage.googleapis.com"
|
||||
}
|
||||
```
|
||||
|
||||
### Timeout IA
|
||||
|
||||
`request_timeout_seconds` vaut `300` par défaut. Si Ollama charge un gros modèle ou répond lentement, augmentez cette valeur. Mettez `0` pour désactiver le timeout côté application.
|
||||
|
||||
### Configuration maths LaTeX
|
||||
|
||||
Pour laisser le modèle générer du LaTeX tout en indiquant clairement les débuts/fins d'équations :
|
||||
|
||||
```bash
|
||||
cp config.math-latex.template config.json
|
||||
```
|
||||
|
||||
Cette config garde :
|
||||
|
||||
```json
|
||||
"math_text_format": "plain"
|
||||
```
|
||||
|
||||
Donc l'application ne transforme rien. Elle tape littéralement la réponse reçue, caractère par caractère.
|
||||
|
||||
Exemple de réponse demandée au modèle :
|
||||
Le prompt système demande au modèle de produire du LaTeX encadré par `[EQ]...[/EQ]`, par exemple :
|
||||
|
||||
```text
|
||||
Les racines sont [EQ]z_1 = x + iy[/EQ] et [EQ]z_2 = x - iy[/EQ].
|
||||
```
|
||||
|
||||
Pour les fractions, intégrales, sommes, etc., le modèle peut utiliser du LaTeX standard dans les balises :
|
||||
À chaque `[EQ]` l'application envoie `Alt+=` (ouvre une équation inline), tape le LaTeX littéralement, puis envoie `→` à chaque `[/EQ]`. Ce comportement est désactivé par défaut sur les autres profils (le texte est tapé tel quel).
|
||||
|
||||
```text
|
||||
On obtient [EQ]\frac{a+b}{c+d}[/EQ] puis [EQ]\int_0^1 f(x)\,dx[/EQ].
|
||||
```
|
||||
## Zone de notification (icône)
|
||||
|
||||
Modes disponibles :
|
||||
L'application se lance sans fenêtre visible. Le menu de l'icône propose :
|
||||
|
||||
- `plain` : mode recommandé ; injecte la réponse exactement telle que le modèle l'a renvoyée.
|
||||
- `unicode` : ancien mode texte Unicode (`z_1` → `z₁`, `x^2` → `x²`) sans objet équation.
|
||||
- **Ouvrir les logs** — fenêtre des journaux en temps réel (également écrits dans `%APPDATA%\ai-typewriter\logs\app.log`).
|
||||
- **Ajouter un profil** — formulaire (nom, fournisseur, modèle, serveur, prompt, équations).
|
||||
- **Modifier le profil** — sous-menu listant tous les profils pour choisir le profil actif.
|
||||
- **Gérer l'authentification** — enregistrer les clés d'API des fournisseurs.
|
||||
- **Quitter** — arrête le processus.
|
||||
|
||||
## Compilation
|
||||
### Sélection et téléchargement des modèles
|
||||
|
||||
### Windows
|
||||
Dans le formulaire de profil, « Choisir / télécharger… » ouvre un sélecteur qui :
|
||||
|
||||
Après clonage du dépôt, lancez simplement :
|
||||
- liste automatiquement les modèles déjà disponibles localement (Ollama `/api/tags`) ;
|
||||
- si connecté à Internet, permet de rechercher dans la bibliothèque publique d'Ollama, de vérifier un modèle exact et de lancer son téléchargement (`ollama pull`).
|
||||
|
||||
### Authentification des fournisseurs
|
||||
|
||||
Les clés d'API ne sont **jamais écrites** dans le fichier de configuration. Chaque profil référence une clé par un nom logique ; la clé est stockée de façon sécurisée dans le **Gestionnaire d'identifiants de Windows** (via `keyring`). Le menu **Gérer l'authentification** permet de les enregistrer, vérifier ou supprimer.
|
||||
|
||||
## Compilation (Windows)
|
||||
|
||||
```bat
|
||||
build.bat
|
||||
```
|
||||
|
||||
Le script crée `.venv`, installe les dépendances, nettoie les anciens artefacts puis génère un exécutable Windows autonome :
|
||||
Le script crée `.venv`, installe les dépendances, puis produit un exécutable autonome **sans console** dans `dist\ai-typewriter.exe`. Il tourne directement en zone de notification.
|
||||
|
||||
```text
|
||||
dist\ai-typewriter.exe
|
||||
```
|
||||
|
||||
### Commande PyInstaller équivalente
|
||||
## Test rapide sans hook clavier ni icône
|
||||
|
||||
```bash
|
||||
pyinstaller --onefile --paths src --name ai-typewriter.exe main.py
|
||||
python main.py --debug --ask "Résume: bonjour tout le monde"
|
||||
```
|
||||
|
||||
L'exécutable est généré dans `dist/`. Le binaire n'est pas versionné Git.
|
||||
Envoie le prompt au profil actif et imprime la réponse brute (aucune icône ni interception clavier).
|
||||
|
||||
## Test rapide sans hook clavier
|
||||
## Tests
|
||||
|
||||
```bash
|
||||
python main.py --config config.json --ask "Résume: bonjour tout le monde"
|
||||
. .venv/bin/activate
|
||||
pytest
|
||||
```
|
||||
|
||||
La suite couvre le découpage en actions (caractères/équations), le stepper, le dépôt de profils, les clients IA (Ollama/Gemini/OpenAI), le stockage sécurisé (keyring mocké) et le catalogue de modèles — sans réel hook clavier, réseau ni Gestionnaire d'identifiants.
|
||||
@@ -1,6 +1,8 @@
|
||||
@echo off
|
||||
setlocal
|
||||
|
||||
REM ============================================================
|
||||
REM Build Windows one-file executable (icône zone de notification)
|
||||
REM ============================================================
|
||||
cd /d "%~dp0"
|
||||
|
||||
if not exist ".venv\Scripts\python.exe" (
|
||||
@@ -18,10 +20,19 @@ if errorlevel 1 goto :error
|
||||
echo [3/4] Nettoyage des anciens builds...
|
||||
if exist build rmdir /s /q build
|
||||
if exist dist rmdir /s /q dist
|
||||
if exist ai-typewriter.exe.spec del /q ai-typewriter.exe.spec
|
||||
if exist ai-typewriter.spec del /q ai-typewriter.spec
|
||||
|
||||
echo [4/4] Compilation Windows one-file...
|
||||
".venv\Scripts\python.exe" -m PyInstaller --clean --onefile --paths src --name ai-typewriter.exe main.py
|
||||
echo [4/4] Compilation Windows one-file (sans console)...
|
||||
".venv\Scripts\python.exe" -m PyInstaller ^
|
||||
--clean --onefile --windowed ^
|
||||
--name ai-typewriter ^
|
||||
--hidden-import keyring ^
|
||||
--hidden-import keyring.backends ^
|
||||
--hidden-import keyring.backends.Windows ^
|
||||
--hidden-import pystray._win32 ^
|
||||
--hidden-import PIL._tkinter_finder ^
|
||||
--paths src ^
|
||||
main.py
|
||||
if errorlevel 1 goto :error
|
||||
|
||||
echo.
|
||||
@@ -31,4 +42,4 @@ exit /b 0
|
||||
:error
|
||||
echo.
|
||||
echo Build echoue.
|
||||
exit /b 1
|
||||
exit /b 1
|
||||
@@ -1,11 +0,0 @@
|
||||
{
|
||||
"provider": "ollama",
|
||||
"model": "llama3.1",
|
||||
"api_key": "",
|
||||
"server_url": "http://localhost:11434",
|
||||
"hotkey": "ctrl+alt+a",
|
||||
"request_timeout_seconds": 300,
|
||||
"math_text_format": "plain",
|
||||
"type_delay_seconds": 0,
|
||||
"system_prompt": "Réponds directement et de manière ultra-concise. Aucune phrase d'introduction, aucune salutation, aucun formatage superflu. Uniquement la réponse brute."
|
||||
}
|
||||
@@ -1,11 +0,0 @@
|
||||
{
|
||||
"provider": "ollama",
|
||||
"model": "llama3.1",
|
||||
"api_key": "",
|
||||
"server_url": "http://localhost:11434",
|
||||
"hotkey": "ctrl+alt+a",
|
||||
"request_timeout_seconds": 300,
|
||||
"math_text_format": "plain",
|
||||
"type_delay_seconds": 0,
|
||||
"system_prompt": "Tu réponds directement, sans salutation ni introduction. Rédige une réponse claire, correcte et concise. Pour toute expression mathématique, formule, calcul, égalité, fraction, somme, intégrale, matrice ou symbole qui doit être traité comme une équation, encadre exactement le bloc avec [EQ] au début et [/EQ] à la fin. Dans ces blocs, écris du LaTeX standard, car il est plus simple et fiable à générer : \frac{a}{b}, z_1, x^2, \int_0^1, \sum_{k=1}^n, etc. N'utilise pas de délimiteurs LaTeX supplémentaires dans les blocs : pas de $, $$, \\(, \\[. Le texte hors des balises [EQ]...[/EQ] reste du texte normal. Exemple valide : Les racines sont [EQ]z_1 = x + iy[/EQ] et [EQ]z_2 = x - iy[/EQ]."
|
||||
}
|
||||
+9
-3
@@ -4,17 +4,23 @@ build-backend = "setuptools.build_meta"
|
||||
|
||||
[project]
|
||||
name = "ai-typewriter"
|
||||
version = "0.1.0"
|
||||
description = "Global hotkey AI key-stepper"
|
||||
version = "0.2.0"
|
||||
description = "AI-typewriter : application d'arrière-plan (icône zone de notification), modèle IA, dactylographie par touches + équations LaTeX"
|
||||
requires-python = ">=3.10"
|
||||
dependencies = [
|
||||
"keyboard==0.13.5",
|
||||
"pyperclip==1.9.0",
|
||||
"requests==2.32.5",
|
||||
"keyring>=25.0",
|
||||
"pystray>=0.19",
|
||||
"Pillow>=10.0",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
dev = ["pytest>=8.0"]
|
||||
|
||||
[tool.setuptools.packages.find]
|
||||
where = ["src"]
|
||||
|
||||
[tool.pytest.ini_options]
|
||||
pythonpath = ["src"]
|
||||
pythonpath = ["src"]
|
||||
+4
-1
@@ -1,5 +1,8 @@
|
||||
keyboard==0.13.5
|
||||
pyperclip==1.9.0
|
||||
requests==2.32.5
|
||||
keyring>=25.0
|
||||
pystray>=0.19
|
||||
Pillow>=10.0
|
||||
pyinstaller==6.16.0
|
||||
pytest==8.4.2
|
||||
pytest==8.4.2
|
||||
@@ -1,3 +1,3 @@
|
||||
"""AI Typewriter package."""
|
||||
"""AI-Typewriter : application d'arrière-plan (icône zone de notification)."""
|
||||
|
||||
__version__ = "0.1.0"
|
||||
__version__ = "0.2.0"
|
||||
+124
-33
@@ -1,50 +1,82 @@
|
||||
"""Clients IA : Ollama, Gemini et tous les fournisseurs OpenAI-compatibles.
|
||||
|
||||
Les clés d'API sont résolues via le gestionnaire de références fourni
|
||||
(`resolve_key`) et ne sont jamais consignées dans les journaux.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Callable
|
||||
|
||||
import requests
|
||||
|
||||
from .config import AppConfig
|
||||
from .config import Profile
|
||||
from .credentials import SecureStore, get_cred
|
||||
|
||||
|
||||
class AIClientError(RuntimeError):
|
||||
"""Raised when the configured AI backend cannot return text."""
|
||||
pass
|
||||
|
||||
|
||||
def ask_ai(prompt: str, config: AppConfig) -> str:
|
||||
ProviderResolver = Callable[[Profile, SecureStore], str]
|
||||
|
||||
|
||||
def _default_resolver(profile: Profile, store: SecureStore) -> str:
|
||||
return get_cred(store, profile.credential)
|
||||
|
||||
|
||||
def ask_ai(
|
||||
prompt: str,
|
||||
profile: Profile,
|
||||
store: SecureStore | None = None,
|
||||
resolve_key: ProviderResolver = _default_resolver,
|
||||
) -> str:
|
||||
"""Envoie `prompt` au modèle du profil et retourne la réponse brute."""
|
||||
if not prompt.strip():
|
||||
raise AIClientError("Le texte capturé est vide.")
|
||||
|
||||
if config.provider == "ollama":
|
||||
return _ask_ollama(prompt, config)
|
||||
if config.provider == "gemini":
|
||||
return _ask_gemini(prompt, config)
|
||||
raise AIClientError(f"Provider non supporté: {config.provider}")
|
||||
provider = profile.provider.lower()
|
||||
if provider == "ollama":
|
||||
return _ask_ollama(prompt, profile)
|
||||
if provider in ("openai", "openrouter", "custom"):
|
||||
return _ask_openai(prompt, profile, store, resolve_key)
|
||||
if provider == "gemini":
|
||||
return _ask_gemini(prompt, profile, store, resolve_key)
|
||||
raise AIClientError(f"Fournisseur non supporté : {profile.provider}")
|
||||
|
||||
|
||||
def _ask_ollama(prompt: str, config: AppConfig) -> str:
|
||||
url = f"{config.server_url}/api/chat"
|
||||
def _timeout(profile: Profile) -> float | None:
|
||||
return profile.request_timeout_seconds
|
||||
|
||||
|
||||
def _timeout_label(profile: Profile) -> str:
|
||||
t = profile.request_timeout_seconds
|
||||
return "désactivé" if t is None else f"{t:.0f} s"
|
||||
|
||||
|
||||
def _ask_ollama(prompt: str, profile: Profile) -> str:
|
||||
url = f"{profile.server_url}/api/chat"
|
||||
payload = {
|
||||
"model": config.model,
|
||||
"model": profile.model,
|
||||
"stream": False,
|
||||
"messages": [
|
||||
{"role": "system", "content": config.system_prompt},
|
||||
{"role": "system", "content": profile.effective_prompt()},
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
}
|
||||
try:
|
||||
response = requests.post(url, json=payload, timeout=config.request_timeout_seconds)
|
||||
response = requests.post(url, json=payload, timeout=_timeout(profile))
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
except requests.Timeout as exc:
|
||||
timeout_label = "désactivé" if config.request_timeout_seconds is None else f"{config.request_timeout_seconds:.0f} s"
|
||||
raise AIClientError(
|
||||
"Ollama n'a pas répondu avant le délai configuré "
|
||||
f"({timeout_label}). Le modèle est peut-être en chargement ou trop lent; "
|
||||
"augmentez request_timeout_seconds dans config.json, ou mettez 0 pour désactiver le timeout."
|
||||
f"({_timeout_label(profile)}). Augmentez request_timeout_seconds, "
|
||||
"ou mettez 0 pour désactiver le timeout."
|
||||
) from exc
|
||||
except requests.RequestException as exc:
|
||||
raise AIClientError(f"Erreur Ollama: {exc}") from exc
|
||||
raise AIClientError(f"Erreur Ollama : {exc}") from exc
|
||||
except ValueError as exc:
|
||||
raise AIClientError("Réponse Ollama invalide: JSON illisible") from exc
|
||||
raise AIClientError("Réponse Ollama invalide : JSON illisible") from exc
|
||||
|
||||
content = data.get("message", {}).get("content")
|
||||
if not isinstance(content, str) or not content.strip():
|
||||
@@ -52,37 +84,96 @@ def _ask_ollama(prompt: str, config: AppConfig) -> str:
|
||||
return content.strip()
|
||||
|
||||
|
||||
def _ask_gemini(prompt: str, config: AppConfig) -> str:
|
||||
if not config.api_key:
|
||||
raise AIClientError("api_key est obligatoire pour provider='gemini'.")
|
||||
def _resolve_key(profile: Profile, store: SecureStore | None, resolve: ProviderResolver) -> str:
|
||||
if store is None:
|
||||
store = SecureStore()
|
||||
key = resolve(profile, store)
|
||||
if not key:
|
||||
raise AIClientError(
|
||||
f"Aucune clé d'API configurée pour le profil « {profile.name} ». "
|
||||
"Ajoutez une authentification pour le fournisseur via l'icône de l'application."
|
||||
)
|
||||
return key
|
||||
|
||||
base = config.server_url or "https://generativelanguage.googleapis.com"
|
||||
url = f"{base}/v1beta/models/{config.model}:generateContent"
|
||||
|
||||
def _ask_openai(
|
||||
prompt: str,
|
||||
profile: Profile,
|
||||
store: SecureStore | None,
|
||||
resolve_key: ProviderResolver,
|
||||
) -> str:
|
||||
"""Appels OpenAI-compatibles (OpenAI, OpenRouter, LM Studio, etc.)."""
|
||||
key = _resolve_key(profile, store, resolve_key)
|
||||
base = (profile.server_url or "https://api.openai.com/v1").rstrip("/")
|
||||
url = f"{base}/chat/completions"
|
||||
headers = {
|
||||
"Authorization": f"Bearer {key}",
|
||||
"Content-Type": "application/json",
|
||||
}
|
||||
payload = {
|
||||
"systemInstruction": {"parts": [{"text": config.system_prompt}]},
|
||||
"model": profile.model,
|
||||
"messages": [
|
||||
{"role": "system", "content": profile.effective_prompt()},
|
||||
{"role": "user", "content": prompt},
|
||||
],
|
||||
}
|
||||
try:
|
||||
response = requests.post(
|
||||
url, json=payload, headers=headers, timeout=_timeout(profile)
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
except requests.Timeout as exc:
|
||||
raise AIClientError(
|
||||
"Le fournisseur n'a pas répondu avant le délai configuré "
|
||||
f"({_timeout_label(profile)}). Augmentez request_timeout_seconds."
|
||||
) from exc
|
||||
except requests.RequestException as exc:
|
||||
raise AIClientError(f"Erreur {profile.provider} : {exc}") from exc
|
||||
except ValueError as exc:
|
||||
raise AIClientError("Réponse du fournisseur invalide : JSON illisible") from exc
|
||||
|
||||
try:
|
||||
text = data["choices"][0]["message"]["content"]
|
||||
except (KeyError, IndexError, TypeError) as exc:
|
||||
raise AIClientError("Réponse du fournisseur vide ou inattendue") from exc
|
||||
if not isinstance(text, str) or not text.strip():
|
||||
raise AIClientError("Réponse vide")
|
||||
return text.strip()
|
||||
|
||||
|
||||
def _ask_gemini(
|
||||
prompt: str,
|
||||
profile: Profile,
|
||||
store: SecureStore | None,
|
||||
resolve_key: ProviderResolver,
|
||||
) -> str:
|
||||
key = _resolve_key(profile, store, resolve_key)
|
||||
base = profile.server_url or "https://generativelanguage.googleapis.com"
|
||||
url = f"{base}/v1beta/models/{profile.model}:generateContent"
|
||||
payload = {
|
||||
"systemInstruction": {"parts": [{"text": profile.effective_prompt()}]},
|
||||
"contents": [{"role": "user", "parts": [{"text": prompt}]}],
|
||||
"generationConfig": {"temperature": 0.2},
|
||||
}
|
||||
try:
|
||||
response = requests.post(
|
||||
url,
|
||||
params={"key": config.api_key},
|
||||
params={"key": key},
|
||||
json=payload,
|
||||
timeout=config.request_timeout_seconds,
|
||||
timeout=_timeout(profile),
|
||||
)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
except requests.Timeout as exc:
|
||||
timeout_label = "désactivé" if config.request_timeout_seconds is None else f"{config.request_timeout_seconds:.0f} s"
|
||||
raise AIClientError(
|
||||
"Gemini n'a pas répondu avant le délai configuré "
|
||||
f"({timeout_label}). Augmentez request_timeout_seconds dans config.json, "
|
||||
"ou mettez 0 pour désactiver le timeout."
|
||||
f"({_timeout_label(profile)}). Augmentez request_timeout_seconds."
|
||||
) from exc
|
||||
except requests.RequestException as exc:
|
||||
raise AIClientError(f"Erreur Gemini: {exc}") from exc
|
||||
raise AIClientError(f"Erreur Gemini : {exc}") from exc
|
||||
except ValueError as exc:
|
||||
raise AIClientError("Réponse Gemini invalide: JSON illisible") from exc
|
||||
raise AIClientError("Réponse Gemini invalide : JSON illisible") from exc
|
||||
|
||||
try:
|
||||
text = data["candidates"][0]["content"]["parts"][0]["text"]
|
||||
@@ -90,4 +181,4 @@ def _ask_gemini(prompt: str, config: AppConfig) -> str:
|
||||
raise AIClientError("Réponse Gemini vide ou inattendue") from exc
|
||||
if not isinstance(text, str) or not text.strip():
|
||||
raise AIClientError("Réponse Gemini vide")
|
||||
return text.strip()
|
||||
return text.strip()
|
||||
+55
-75
@@ -1,107 +1,87 @@
|
||||
"""Point d'entrée de l'application.
|
||||
|
||||
Par défaut, l'application s'exécute en arrière-plan, sans fenêtre visible,
|
||||
avec une icône dans la zone de notification (Windows). Deux modes console
|
||||
restent disponibles pour le développement / le test :
|
||||
- `python main.py --ask "texte"` : envoie le texte au profil actif et
|
||||
imprime la réponse sans intercepter le clavier.
|
||||
- `python main.py` : lance l'icône de la zone de notification.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
import logging
|
||||
import sys
|
||||
import time
|
||||
from threading import Lock
|
||||
|
||||
import keyboard
|
||||
|
||||
from .ai_client import AIClientError, ask_ai
|
||||
from .clipboard_capture import capture_clipboard
|
||||
from .config import AppConfig, load_config
|
||||
from .key_stepper import KeyStepper
|
||||
from .math_format import format_math_text
|
||||
from .config import ConfigStore, load_config
|
||||
from .credentials import SecureStore
|
||||
from .engine import AITypewriterEngine
|
||||
from .logging_utils import setup_logging
|
||||
|
||||
LOG = logging.getLogger("ai_typewriter")
|
||||
|
||||
|
||||
def prepare_type_payload(answer: str, mode: str) -> str:
|
||||
return format_math_text(answer, mode)
|
||||
|
||||
|
||||
class AITypewriterApp:
|
||||
def __init__(self, config: AppConfig) -> None:
|
||||
self.config = config
|
||||
self._busy = Lock()
|
||||
self._stepper: KeyStepper | None = None
|
||||
|
||||
def run(self) -> None:
|
||||
keyboard.add_hotkey(self.config.hotkey, self._handle_hotkey, suppress=False)
|
||||
LOG.info("Prêt. Raccourci: %s. Quitter: Ctrl+C dans ce terminal.", self.config.hotkey)
|
||||
keyboard.wait()
|
||||
|
||||
def _handle_hotkey(self) -> None:
|
||||
if not self._busy.acquire(blocking=False):
|
||||
LOG.warning("Requête déjà en cours, raccourci ignoré.")
|
||||
return
|
||||
try:
|
||||
self._capture_ask_and_step()
|
||||
finally:
|
||||
self._busy.release()
|
||||
|
||||
def _capture_ask_and_step(self) -> None:
|
||||
try:
|
||||
selected = capture_clipboard()
|
||||
LOG.info("Texte lu depuis le presse-papier: %d caractères.", len(selected))
|
||||
LOG.info("Texte capturé: %r", selected)
|
||||
started_at = time.perf_counter()
|
||||
answer = ask_ai(selected, self.config)
|
||||
answer_to_type = prepare_type_payload(answer, self.config.math_text_format)
|
||||
elapsed = time.perf_counter() - started_at
|
||||
LOG.info(
|
||||
"Réponse reçue en %.2f s: %d caractères à écrire en mode %s. Appuyez sur une touche pour écrire chaque caractère.",
|
||||
elapsed,
|
||||
len(answer_to_type),
|
||||
self.config.math_text_format,
|
||||
)
|
||||
if self._stepper is not None:
|
||||
self._stepper.stop()
|
||||
self._stepper = KeyStepper(answer_to_type, delay=self.config.type_delay_seconds)
|
||||
self._stepper.start()
|
||||
except (AIClientError, FileNotFoundError, ValueError) as exc:
|
||||
LOG.error("%s", exc)
|
||||
except Exception:
|
||||
LOG.exception("Erreur inattendue")
|
||||
|
||||
|
||||
def build_parser() -> argparse.ArgumentParser:
|
||||
parser = argparse.ArgumentParser(description="AI Typewriter")
|
||||
parser.add_argument("--config", default="config.json", help="Chemin du fichier config.json")
|
||||
parser = argparse.ArgumentParser(description="AI-Typewriter")
|
||||
parser.add_argument(
|
||||
"--config",
|
||||
default=None,
|
||||
help="Chemin du fichier config.json (défaut: dossier appdata)",
|
||||
)
|
||||
parser.add_argument("--debug", action="store_true", help="Logs détaillés")
|
||||
parser.add_argument("--ask", help="Mode test: envoie ce texte à l'IA et imprime la réponse, sans hook clavier")
|
||||
parser.add_argument(
|
||||
"--ask",
|
||||
help="Mode test : envoie ce texte à l'IA du profil actif et imprime "
|
||||
"la réponse, sans hook clavier ni icône.",
|
||||
)
|
||||
return parser
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
args = build_parser().parse_args(argv)
|
||||
logging.basicConfig(
|
||||
level=logging.DEBUG if args.debug else logging.INFO,
|
||||
format="%(asctime)s %(levelname)s %(message)s",
|
||||
)
|
||||
setup_logging(level=logging.DEBUG if args.debug else logging.INFO)
|
||||
|
||||
try:
|
||||
config = load_config(args.config)
|
||||
store = load_config(args.config)
|
||||
except Exception as exc:
|
||||
LOG.error("%s", exc)
|
||||
return 2
|
||||
|
||||
engine = AITypewriterEngine(store, SecureStore())
|
||||
|
||||
if args.ask is not None:
|
||||
try:
|
||||
started_at = time.perf_counter()
|
||||
answer = ask_ai(args.ask, config)
|
||||
answer_to_type = prepare_type_payload(answer, config.math_text_format)
|
||||
elapsed = time.perf_counter() - started_at
|
||||
LOG.info("Réponse générée en %.2f s: %d caractères en mode %s.", elapsed, len(answer_to_type), config.math_text_format)
|
||||
print(answer_to_type)
|
||||
answer = engine.ask_only(args.ask)
|
||||
print(answer)
|
||||
return 0
|
||||
except AIClientError as exc:
|
||||
except Exception as exc:
|
||||
LOG.error("%s", exc)
|
||||
return 1
|
||||
|
||||
AITypewriterApp(config).run()
|
||||
return run_tray(store)
|
||||
|
||||
|
||||
def run_tray(store: ConfigStore) -> int:
|
||||
"""Lance l'application en arrière-plan avec l'icône de notification."""
|
||||
# Import différé pour garantir que --ask / les tests fonctionnent même si
|
||||
# pystray ou PIL ne sont pas installés / sans affichage graphique.
|
||||
from .tray import TrayApp
|
||||
|
||||
app = TrayApp(store)
|
||||
app.start()
|
||||
|
||||
# Maintient le processus en vie ; pystray run() tourne déjà dans un thread.
|
||||
try:
|
||||
# Boucle événementielle tant que l'application n'est pas arrêtée.
|
||||
while True:
|
||||
import time
|
||||
|
||||
time.sleep(3600)
|
||||
except KeyboardInterrupt:
|
||||
app.stop()
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main(sys.argv[1:]))
|
||||
raise SystemExit(main(sys.argv[1:]))
|
||||
+256
-58
@@ -1,71 +1,269 @@
|
||||
"""Profils et configuration de l'application.
|
||||
|
||||
La configuration est stockée dans un fichier JSON, créé par défaut dans
|
||||
le répertoire des données de l'application (Windows : %APPDATA%\\ai-typewriter\\config.json).
|
||||
Elle contient une liste de profils nommés, le profil actif et les réglages
|
||||
globaux (raccourci clavier global, etc.).
|
||||
|
||||
Les clés d'API ne sont jamais écrites dans ce fichier : elles sont stockées
|
||||
dans le gestionnaire de références de Windows (gestionnaire d'identifiants)
|
||||
via la bibliothèque `keyring`. Ce fichier ne référence le secret que par un
|
||||
nom logique (`credential`).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from dataclasses import dataclass
|
||||
import os
|
||||
import threading
|
||||
from dataclasses import dataclass, field, asdict
|
||||
from pathlib import Path
|
||||
from typing import Any, Literal
|
||||
|
||||
Provider = Literal["ollama", "gemini"]
|
||||
MathTextFormat = Literal["plain", "unicode", "unicode_math"]
|
||||
from typing import Any
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class AppConfig:
|
||||
provider: Provider = "ollama"
|
||||
model: str = "llama3.1"
|
||||
api_key: str = ""
|
||||
server_url: str = "http://localhost:11434"
|
||||
hotkey: str = "ctrl+alt+a"
|
||||
request_timeout_seconds: float | None = 300.0
|
||||
copy_wait_seconds: float = 1.0
|
||||
type_delay_seconds: float = 0.0
|
||||
math_text_format: MathTextFormat = "plain"
|
||||
restore_clipboard: bool = True
|
||||
system_prompt: str = (
|
||||
# ---------------------------------------------------------------------------
|
||||
# Chemins & constantes
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
APP_NAME = "ai-typewriter"
|
||||
|
||||
|
||||
def appdata_dir() -> Path:
|
||||
"""Répertoire de données applicatives (Linux : ~/.config/ai-typewriter).
|
||||
|
||||
Sur Windows ce sera %APPDATA%\\ai-typewriter ; sur les autres plateformes
|
||||
on retombe sur le répertoire utilisateur pour rester fonctionnel.
|
||||
"""
|
||||
base = os.environ.get("APPDATA")
|
||||
if base:
|
||||
return Path(base) / APP_NAME
|
||||
homedir = Path.home()
|
||||
if os.name == "nt":
|
||||
return homedir / "AppData" / "Roaming" / APP_NAME
|
||||
return homedir / ".config" / APP_NAME
|
||||
|
||||
|
||||
def config_path() -> Path:
|
||||
return appdata_dir() / "config.json"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Marquage des équations (profil « math »)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
DEFAULT_EQ_START_MARKER = "[EQ]"
|
||||
DEFAULT_EQ_END_MARKER = "[/EQ]"
|
||||
# Séquence de touches envoyée : Alt+= ouvre une équation inline (Word/OneNote).
|
||||
DEFAULT_EQ_START_KEY = "alt+="
|
||||
# Enregistrement Word pour sortir du champ d'équation.
|
||||
DEFAULT_EQ_END_KEY = "right"
|
||||
|
||||
|
||||
def default_system_prompt() -> str:
|
||||
return (
|
||||
"Réponds directement et de manière ultra-concise. "
|
||||
"Aucune phrase d'introduction, aucune salutation, aucun formatage superflu. "
|
||||
"Uniquement la réponse brute."
|
||||
)
|
||||
|
||||
|
||||
def _coerce_provider(value: Any) -> Provider:
|
||||
provider = str(value or "ollama").lower().strip()
|
||||
if provider not in {"ollama", "gemini"}:
|
||||
raise ValueError("config.provider doit être 'ollama' ou 'gemini'")
|
||||
return provider # type: ignore[return-value]
|
||||
|
||||
|
||||
def _coerce_timeout(value: Any) -> float | None:
|
||||
timeout = float(value)
|
||||
if timeout <= 0:
|
||||
return None
|
||||
return timeout
|
||||
|
||||
|
||||
def _coerce_math_text_format(value: Any) -> MathTextFormat:
|
||||
mode = str(value or "plain").lower().strip()
|
||||
if mode not in {"plain", "unicode", "unicode_math"}:
|
||||
raise ValueError("config.math_text_format doit être 'plain' ou 'unicode'")
|
||||
return mode # type: ignore[return-value]
|
||||
|
||||
|
||||
def load_config(path: str | Path = "config.json") -> AppConfig:
|
||||
cfg_path = Path(path)
|
||||
if not cfg_path.exists():
|
||||
raise FileNotFoundError(
|
||||
f"Configuration introuvable: {cfg_path}. Copiez config.json.template vers config.json."
|
||||
)
|
||||
data = json.loads(cfg_path.read_text(encoding="utf-8"))
|
||||
return AppConfig(
|
||||
provider=_coerce_provider(data.get("provider", "ollama")),
|
||||
model=str(data.get("model", "llama3.1")),
|
||||
api_key=str(data.get("api_key", "")),
|
||||
server_url=str(data.get("server_url", "http://localhost:11434")).rstrip("/"),
|
||||
hotkey=str(data.get("hotkey", "ctrl+alt+a")).lower(),
|
||||
request_timeout_seconds=_coerce_timeout(data.get("request_timeout_seconds", 300)),
|
||||
copy_wait_seconds=float(data.get("copy_wait_seconds", 1)),
|
||||
type_delay_seconds=float(data.get("type_delay_seconds", 0)),
|
||||
math_text_format=_coerce_math_text_format(data.get("math_text_format", "plain")),
|
||||
restore_clipboard=bool(data.get("restore_clipboard", True)),
|
||||
system_prompt=str(data.get("system_prompt", AppConfig.system_prompt)),
|
||||
def default_math_latex_prompt() -> str:
|
||||
return (
|
||||
"Tu réponds directement, sans salutation ni introduction. "
|
||||
"Rédige une réponse claire, correcte et concise. "
|
||||
"Pour toute expression mathématique, formule, calcul, égalité, fraction, somme, "
|
||||
"intégrale, matrice ou symbole destiné à être traité comme une équation, encadre "
|
||||
"exactement le bloc avec [EQ] au début et [/EQ] à la fin. Dans ces blocs écris du "
|
||||
"LaTeX standard, plus simple et fiable à générer : \\frac{a}{b}, z_1, x^2, "
|
||||
"\\int_0^1, \\sum_{k=1}^n, etc. N'utilise pas de délimiteurs LaTeX supplémentaires "
|
||||
"dans les blocs (pas de $, $$, \\\\(, \\\\[). Le texte hors des balises "
|
||||
"[EQ]...[/EQ] reste du texte normal. Exemple valide : "
|
||||
"Les racines sont [EQ]z_1 = x + iy[/EQ] et [EQ]z_2 = x - iy[/EQ]."
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Modèle de profil
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class Profile:
|
||||
"""Un profil = un fournisseur + un modèle + un prompt système + des réglages."""
|
||||
|
||||
name: str = "Général"
|
||||
provider: str = "ollama" # ollama | openai | gemini
|
||||
model: str = "llama3.1"
|
||||
server_url: str = "http://localhost:11434"
|
||||
request_timeout_seconds: float | None = 300.0
|
||||
type_delay_seconds: float = 0.0
|
||||
system_prompt: str = ""
|
||||
# Gestion des équations LaTeX : si actif, les marqueurs sont interceptés et
|
||||
# remplacés par des séquences de touches.
|
||||
equation_enabled: bool = False
|
||||
eq_start_marker: str = DEFAULT_EQ_START_MARKER
|
||||
eq_end_marker: str = DEFAULT_EQ_END_MARKER
|
||||
eq_start_key: str = DEFAULT_EQ_START_KEY
|
||||
eq_end_key: str = DEFAULT_EQ_END_KEY
|
||||
# Nom logique de la référence stockée dans le gestionnaire de références
|
||||
# (vide si le fournisseur n'exige pas de clé, ex : Ollama local).
|
||||
credential: str = ""
|
||||
# Réglages de capture.
|
||||
copy_wait_seconds: float = 1.0
|
||||
restore_clipboard: bool = True
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self.system_prompt:
|
||||
self.system_prompt = default_system_prompt()
|
||||
if self.provider == "ollama":
|
||||
self.server_url = (self.server_url or "http://localhost:11434").rstrip("/")
|
||||
|
||||
def effective_prompt(self) -> str:
|
||||
return self.system_prompt or default_system_prompt()
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return asdict(self)
|
||||
|
||||
@classmethod
|
||||
def from_dict(cls, data: dict[str, Any]) -> "Profile":
|
||||
known = {f for f in cls.__dataclass_fields__} # type: ignore[attr-defined]
|
||||
return cls(**{k: v for k, v in data.items() if k in known})
|
||||
|
||||
|
||||
def math_latex_profile(name: str = "Mathématiques (LaTeX)", **overrides: Any) -> Profile:
|
||||
"""Profil d'exemple spécialisé en mathématiques (LaTeX encadré)."""
|
||||
data = dict(
|
||||
name=name,
|
||||
provider="ollama",
|
||||
model="llama3.1",
|
||||
server_url="http://localhost:11434",
|
||||
system_prompt=default_math_latex_prompt(),
|
||||
equation_enabled=True,
|
||||
eq_start_marker=DEFAULT_EQ_START_MARKER,
|
||||
eq_end_marker=DEFAULT_EQ_END_MARKER,
|
||||
eq_start_key=DEFAULT_EQ_START_KEY,
|
||||
eq_end_key=DEFAULT_EQ_END_KEY,
|
||||
)
|
||||
data.update(overrides)
|
||||
return Profile.from_dict(data)
|
||||
|
||||
|
||||
def default_profiles() -> list[Profile]:
|
||||
return [
|
||||
Profile(name="Général"),
|
||||
math_latex_profile(),
|
||||
]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Dépôt de configuration
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
class ConfigError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class ConfigStore:
|
||||
"""Lecture/écriture du fichier de configuration et sélection du profil actif."""
|
||||
|
||||
def __init__(self, path: str | Path | None = None) -> None:
|
||||
self.path = Path(path) if path else config_path()
|
||||
self._lock = threading.Lock()
|
||||
self.active_name: str = "Général"
|
||||
self.profiles: list[Profile] = []
|
||||
self.hotkey: str = "ctrl+alt+a"
|
||||
self._loaded = False
|
||||
|
||||
# -- persistance --------------------------------------------------------
|
||||
|
||||
def ensure_defaults(self) -> None:
|
||||
"""Crée le dossier et le fichier de configuration par défaut si absents."""
|
||||
self.path.parent.mkdir(parents=True, exist_ok=True)
|
||||
if not self.path.exists():
|
||||
self.profiles = default_profiles()
|
||||
self.active_name = self.profiles[0].name
|
||||
self.hotkey = "ctrl+alt+a"
|
||||
self.save()
|
||||
|
||||
def load(self) -> None:
|
||||
self.ensure_defaults()
|
||||
with self._lock:
|
||||
data = json.loads(self.path.read_text(encoding="utf-8"))
|
||||
self.hotkey = str(data.get("hotkey", "ctrl+alt+a")).lower()
|
||||
self.active_name = str(data.get("active_profile", "Général"))
|
||||
raw_profiles = data.get("profiles", [])
|
||||
if not raw_profiles:
|
||||
raw_profiles = [p.to_dict() for p in default_profiles()]
|
||||
self.profiles = [Profile.from_dict(p) for p in raw_profiles]
|
||||
if not self.profiles:
|
||||
raise ConfigError("Aucun profil disponible dans la configuration.")
|
||||
names = [p.name for p in self.profiles]
|
||||
if self.active_name not in names:
|
||||
self.active_name = names[0]
|
||||
self._loaded = True
|
||||
|
||||
def save(self) -> None:
|
||||
with self._lock:
|
||||
body = {
|
||||
"hotkey": self.hotkey,
|
||||
"active_profile": self.active_name,
|
||||
"profiles": [p.to_dict() for p in self.profiles],
|
||||
}
|
||||
self.path.parent.mkdir(parents=True, exist_ok=True)
|
||||
self.path.write_text(
|
||||
json.dumps(body, ensure_ascii=False, indent=2), encoding="utf-8"
|
||||
)
|
||||
|
||||
# -- profils -------------------------------------------------------------
|
||||
|
||||
def get_all(self) -> list[Profile]:
|
||||
if not self._loaded:
|
||||
self.load()
|
||||
return list(self.profiles)
|
||||
|
||||
def get(self, name: str) -> Profile:
|
||||
for p in self.get_all():
|
||||
if p.name == name:
|
||||
return p
|
||||
raise KeyError(name)
|
||||
|
||||
def active(self) -> Profile:
|
||||
if not self._loaded:
|
||||
self.load()
|
||||
for p in self.profiles:
|
||||
if p.name == self.active_name:
|
||||
return p
|
||||
return self.profiles[0]
|
||||
|
||||
def set_active(self, name: str) -> None:
|
||||
self.get(name) # valide l'existence
|
||||
self.active_name = name
|
||||
self.save()
|
||||
|
||||
def upsert(self, profile: Profile) -> None:
|
||||
if not self._loaded:
|
||||
self.load()
|
||||
for i, p in enumerate(self.profiles):
|
||||
if p.name == profile.name:
|
||||
self.profiles[i] = profile
|
||||
break
|
||||
else:
|
||||
self.profiles.append(profile)
|
||||
self.save()
|
||||
|
||||
def remove(self, name: str) -> None:
|
||||
if not self._loaded:
|
||||
self.load()
|
||||
if len(self.profiles) == 1:
|
||||
raise ConfigError("Impossible de supprimer le dernier profil.")
|
||||
self.profiles = [p for p in self.profiles if p.name != name]
|
||||
if self.active_name == name:
|
||||
self.active_name = self.profiles[0].name
|
||||
self.save()
|
||||
|
||||
|
||||
def load_config(path: str | Path | None = None) -> ConfigStore:
|
||||
store = ConfigStore(path)
|
||||
store.load()
|
||||
return store
|
||||
@@ -0,0 +1,86 @@
|
||||
"""Stockage sécurisé des clés d'API.
|
||||
|
||||
Sous Windows, les secrets sont conservés dans le Gestionnaire d'identifiants
|
||||
(Credential Manager) via `keyring`, et jamais écrits en clair dans un fichier
|
||||
de configuration. Sur les autres plateformes, on retombe sur le trousseau
|
||||
fourni par keyring (éventuellement crypté) pour rester utilisable.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
|
||||
LOG = logging.getLogger("ai_typewriter.credentials")
|
||||
|
||||
try: # pragma: no cover - dépend de la plateforme
|
||||
import keyring
|
||||
except Exception: # pragma: no cover
|
||||
keyring = None # type: ignore
|
||||
|
||||
SERVICE = "ai-typewriter"
|
||||
|
||||
|
||||
class CredentialError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
class SecureStore:
|
||||
"""Interface vers le stockage sécurisé des références par fournisseur.
|
||||
|
||||
Une « référence » est identifiée par un nom logique (celui porté par le
|
||||
profil dans `Profile.credential`). Le même nom peut être partagé par
|
||||
plusieurs profils d'un même fournisseur.
|
||||
"""
|
||||
|
||||
def __init__(self, service: str = SERVICE) -> None:
|
||||
self.service = service
|
||||
|
||||
def store(self, credential: str, secret: str) -> None:
|
||||
if not credential:
|
||||
raise CredentialError("Nom de référence invalide.")
|
||||
if keyring is None: # pragma: no cover
|
||||
raise CredentialError(
|
||||
"keyring indisponible : impossible de stocker une clé de façon sécurisée."
|
||||
)
|
||||
try:
|
||||
keyring.set_password(self.service, credential, secret)
|
||||
except Exception as exc: # pragma: no cover
|
||||
LOG.exception("Échec de l'enregistrement de la référence %s", credential)
|
||||
raise CredentialError(
|
||||
f"Impossible d'enregistrer la référence « {credential} ». "
|
||||
f"Vérifiez que le Gestionnaire d'identifiants est disponible : {exc}"
|
||||
) from exc
|
||||
|
||||
def get(self, credential: str) -> str:
|
||||
if not credential:
|
||||
return ""
|
||||
if keyring is None: # pragma: no cover
|
||||
return ""
|
||||
try:
|
||||
value = keyring.get_password(self.service, credential)
|
||||
except Exception: # pragma: no cover
|
||||
LOG.exception("Lecture impossible de la référence %s", credential)
|
||||
return ""
|
||||
return value or ""
|
||||
|
||||
def delete(self, credential: str) -> None:
|
||||
if keyring is None: # pragma: no cover
|
||||
raise CredentialError("keyring indisponible.")
|
||||
try:
|
||||
keyring.delete_password(self.service, credential)
|
||||
except keyring.errors.PasswordDeleteError: # type: ignore[attr-defined]
|
||||
return
|
||||
except Exception as exc: # pragma: no cover
|
||||
raise CredentialError(f"Suppression impossible de « {credential} »") from exc
|
||||
|
||||
def has(self, credential: str) -> bool:
|
||||
return bool(self.get(credential))
|
||||
|
||||
|
||||
def get_cred(store: SecureStore | None, credential: str) -> str:
|
||||
"""Raccourci : lit la référence via le store, en créant un si nécessaire."""
|
||||
if store is None:
|
||||
store = SecureStore()
|
||||
if not credential:
|
||||
return ""
|
||||
return store.get(credential)
|
||||
@@ -0,0 +1,108 @@
|
||||
"""Moteur principal : capture presse-papier -> IA -> dactylographie.
|
||||
|
||||
Ce module est indépendant de l'interface (tray ou console) et peut être
|
||||
réutilisé par l'application graphique, la ligne de commande ou les tests.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
import time
|
||||
|
||||
import keyboard
|
||||
|
||||
from .ai_client import AIClientError, ask_ai
|
||||
from .clipboard_capture import capture_clipboard
|
||||
from .config import ConfigStore, Profile
|
||||
from .credentials import SecureStore
|
||||
from .key_stepper import KeyStepper
|
||||
|
||||
LOG = logging.getLogger("ai_typewriter")
|
||||
|
||||
|
||||
def prepare_type_payload(answer: str, profile: Profile) -> str:
|
||||
"""La réponse est écrite telle quelle (le découpage des marqueurs équation
|
||||
est géré dans KeyStepper). Aucune normalisation Unicode n'est appliquée."""
|
||||
return answer
|
||||
|
||||
|
||||
class AITypewriterEngine:
|
||||
def __init__(self, store: ConfigStore, secure: SecureStore | None = None) -> None:
|
||||
self.store = store
|
||||
self.secure = secure or SecureStore()
|
||||
self._busy = threading.Lock()
|
||||
self._stepper: KeyStepper | None = None
|
||||
|
||||
# -- boucle d'interaction ---------------------------------------------
|
||||
|
||||
def handle_hotkey(self) -> None:
|
||||
if not self._busy.acquire(blocking=False):
|
||||
LOG.warning("Requête déjà en cours, raccourci ignoré.")
|
||||
return
|
||||
try:
|
||||
self.capture_ask_and_step()
|
||||
finally:
|
||||
self._busy.release()
|
||||
|
||||
def capture_ask_and_step(self) -> None:
|
||||
try:
|
||||
selected = capture_clipboard()
|
||||
if not selected.strip():
|
||||
LOG.warning("Presse-papier vide ; rien à envoyer.")
|
||||
return
|
||||
LOG.info("Texte lu depuis le presse-papier : %d caractères.", len(selected))
|
||||
LOG.info("Texte capturé : %r", selected)
|
||||
profile = self.store.active()
|
||||
|
||||
started_at = time.perf_counter()
|
||||
answer = ask_ai(selected, profile, store=self.secure)
|
||||
answer_to_type = prepare_type_payload(answer, profile)
|
||||
elapsed = time.perf_counter() - started_at
|
||||
|
||||
if profile.equation_enabled:
|
||||
mode = f"équations {profile.model}"
|
||||
else:
|
||||
mode = "texte"
|
||||
LOG.info(
|
||||
"Réponse reçue en %.2f s : %d caractères à écrire (%s). "
|
||||
"Appuyez sur une touche pour écrire chaque élément.",
|
||||
elapsed,
|
||||
len(answer_to_type),
|
||||
mode,
|
||||
)
|
||||
|
||||
if self._stepper is not None:
|
||||
self._stepper.stop()
|
||||
self._stepper = KeyStepper(
|
||||
answer_to_type,
|
||||
delay=profile.type_delay_seconds,
|
||||
equation_enabled=profile.equation_enabled,
|
||||
eq_start_marker=profile.eq_start_marker,
|
||||
eq_end_marker=profile.eq_end_marker,
|
||||
eq_start_key=profile.eq_start_key,
|
||||
eq_end_key=profile.eq_end_key,
|
||||
)
|
||||
self._stepper.start()
|
||||
except (AIClientError, FileNotFoundError, ValueError, KeyError) as exc:
|
||||
LOG.error("%s", exc)
|
||||
except Exception:
|
||||
LOG.exception("Erreur inattendue")
|
||||
|
||||
# -- interface publique pour l'UI ---------------------------------------
|
||||
|
||||
def ask_only(self, prompt: str) -> str:
|
||||
"""Envoie un prompt au profil actif et retourne la réponse (sans taper)."""
|
||||
profile = self.store.active()
|
||||
return ask_ai(prompt, profile, store=self.secure)
|
||||
|
||||
def current_profile(self) -> Profile:
|
||||
return self.store.active()
|
||||
|
||||
def stop_stepper(self) -> None:
|
||||
if self._stepper is not None:
|
||||
self._stepper.stop()
|
||||
|
||||
|
||||
def bind_hotkey(engine: AITypewriterEngine, hotkey: str) -> None:
|
||||
keyboard.add_hotkey(hotkey, engine.handle_hotkey, suppress=False)
|
||||
@@ -1,48 +1,132 @@
|
||||
"""Dactylographe : chaque pression de touche physique fait avancer l'écriture.
|
||||
|
||||
Le texte de réponse est découpé en « actions » : caractères littéraux à taper
|
||||
ou séquences de touches à envoyer. Quand la gestion d'équations est active,
|
||||
les marqueurs de début/fin d'équation (ex. `[EQ]` / `[/EQ]`) sont interceptés
|
||||
et remplacés par une séquence de touches (par défaut `Alt+=` pour ouvrir une
|
||||
équation Word et `→` pour en sortir) au lieu d'être tapés littéralement.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from dataclasses import dataclass, field
|
||||
from threading import Lock
|
||||
from typing import Optional
|
||||
from typing import Iterable, Optional
|
||||
|
||||
import keyboard
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Découpage en actions (pure, unité testable sans clavier)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Action:
|
||||
kind: str # "char" | "seq"
|
||||
value: str
|
||||
|
||||
|
||||
_MARKER_TOKENS = object() # injection simple pour les tests
|
||||
|
||||
|
||||
def build_actions(
|
||||
text: str,
|
||||
equation_enabled: bool = False,
|
||||
eq_start_marker: str = "[EQ]",
|
||||
eq_end_marker: str = "[/EQ]",
|
||||
eq_start_key: str = "alt+=",
|
||||
eq_end_key: str = "right",
|
||||
) -> list[Action]:
|
||||
"""Convertit la réponse IA en liste d'actions à exécuter séquentiellement.
|
||||
|
||||
- Sans équation : une action `char` par caractère (comportement d'origine).
|
||||
- Avec équation : les marqueurs sont retirés et remplacés par une action
|
||||
`seq` envoyant la séquence de touches correspondante. Le contenu entre
|
||||
les marqueurs reste tapé caractère par caractère.
|
||||
"""
|
||||
if not equation_enabled or not eq_start_marker or not eq_end_marker:
|
||||
return [Action("char", ch) for ch in text]
|
||||
|
||||
start_re = re.escape(eq_start_marker)
|
||||
end_re = re.escape(eq_end_marker)
|
||||
pattern = re.compile(f"({start_re}|{end_re})")
|
||||
actions: list[Action] = []
|
||||
for part in pattern.split(text):
|
||||
if not part:
|
||||
continue
|
||||
if part == eq_start_marker:
|
||||
actions.append(Action("seq", eq_start_key))
|
||||
elif part == eq_end_marker:
|
||||
actions.append(Action("seq", eq_end_key))
|
||||
else:
|
||||
actions.extend(Action("char", ch) for ch in part)
|
||||
return actions
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Stepper (nécessite le hook clavier ; testé par monkeypatch)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
@dataclass
|
||||
class KeyStepper:
|
||||
text: str
|
||||
delay: float = 0.0
|
||||
equation_enabled: bool = False
|
||||
eq_start_marker: str = "[EQ]"
|
||||
eq_end_marker: str = "[/EQ]"
|
||||
eq_start_key: str = "alt+="
|
||||
eq_end_key: str = "right"
|
||||
_actions: list[Action] = field(default_factory=list)
|
||||
_index: int = 0
|
||||
_hook: Optional[object] = None
|
||||
_lock: Lock = field(default_factory=Lock)
|
||||
_injecting: bool = False
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if not self._actions:
|
||||
self._actions = build_actions(
|
||||
self.text,
|
||||
equation_enabled=self.equation_enabled,
|
||||
eq_start_marker=self.eq_start_marker,
|
||||
eq_end_marker=self.eq_end_marker,
|
||||
eq_start_key=self.eq_start_key,
|
||||
eq_end_key=self.eq_end_key,
|
||||
)
|
||||
|
||||
def start(self) -> None:
|
||||
if not self.text:
|
||||
if not self._actions:
|
||||
return
|
||||
if self._hook is not None:
|
||||
return
|
||||
# suppress=True blocks the physical key while the callback types the next AI character.
|
||||
# suppress=True bloque la touche physique pendant que le callback émet
|
||||
# la prochaine action de la réponse IA.
|
||||
self._hook = keyboard.hook(self._on_event, suppress=True)
|
||||
|
||||
@property
|
||||
def remaining_characters(self) -> int:
|
||||
return max(len(self.text) - self._index, 0)
|
||||
return max(len(self._actions) - self._index, 0)
|
||||
|
||||
def type_next_character(self) -> None:
|
||||
def type_next_action(self) -> None:
|
||||
with self._lock:
|
||||
if self._index >= len(self.text):
|
||||
if self._index >= len(self._actions):
|
||||
self.stop()
|
||||
return
|
||||
char = self.text[self._index]
|
||||
action = self._actions[self._index]
|
||||
self._index += 1
|
||||
|
||||
self._injecting = True
|
||||
try:
|
||||
self._type_char(char)
|
||||
if action.kind == "seq":
|
||||
keyboard.send(action.value)
|
||||
else:
|
||||
self._type_char(action.value)
|
||||
finally:
|
||||
self._injecting = False
|
||||
|
||||
if self._index >= len(self.text):
|
||||
if self._index >= len(self._actions):
|
||||
self.stop()
|
||||
|
||||
def stop(self) -> None:
|
||||
@@ -54,7 +138,7 @@ class KeyStepper:
|
||||
def _on_event(self, event: keyboard.KeyboardEvent) -> None:
|
||||
if self._injecting or event.event_type != keyboard.KEY_DOWN:
|
||||
return
|
||||
self.type_next_character()
|
||||
self.type_next_action()
|
||||
|
||||
def _type_char(self, char: str) -> None:
|
||||
if char == "\n":
|
||||
@@ -62,4 +146,4 @@ class KeyStepper:
|
||||
elif char == "\t":
|
||||
keyboard.send("tab")
|
||||
else:
|
||||
keyboard.write(char, delay=self.delay, exact=True)
|
||||
keyboard.write(char, delay=self.delay, exact=True)
|
||||
@@ -0,0 +1,88 @@
|
||||
"""Journalisation : écriture sur fichier + diffusion vers les vues en direct.
|
||||
|
||||
La fenêtre « Ouvrir les logs » s'abonne à ce module et reçoit les messages en
|
||||
temps réel ; les logs sont également écrits dans %APPDATA%\\ai-typewriter\\logs\\.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import logging.handlers
|
||||
import queue
|
||||
import threading
|
||||
from pathlib import Path
|
||||
|
||||
from .config import appdata_dir
|
||||
|
||||
|
||||
class QueueHandler(logging.Handler):
|
||||
"""Forwarde chaque enregistrement vers une `queue.Queue`."""
|
||||
|
||||
def __init__(self, q: "queue.Queue[logging.LogRecord] | None" = None) -> None:
|
||||
super().__init__()
|
||||
self.queue: queue.Queue = q if q is not None else queue.Queue()
|
||||
|
||||
def emit(self, record: logging.LogRecord) -> None:
|
||||
try:
|
||||
self.queue.put_nowait(record)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
class LogBroadcaster:
|
||||
"""Point central : une file de diffusion + capacité à ajouter des vues."""
|
||||
|
||||
def __init__(self) -> None:
|
||||
self._q: queue.Queue = queue.Queue()
|
||||
self.handler = QueueHandler(self._q)
|
||||
|
||||
def install(self, level: int = logging.INFO) -> None:
|
||||
root = logging.getLogger()
|
||||
root.addHandler(self.handler)
|
||||
root.setLevel(level)
|
||||
self.handler.setLevel(level)
|
||||
|
||||
def drain(self) -> list[logging.LogRecord]:
|
||||
out: list[logging.LogRecord] = []
|
||||
while True:
|
||||
try:
|
||||
out.append(self._q.get_nowait())
|
||||
except queue.Empty:
|
||||
return out
|
||||
|
||||
|
||||
_broadcaster: LogBroadcaster | None = None
|
||||
|
||||
|
||||
def broadcaster() -> LogBroadcaster:
|
||||
global _broadcaster
|
||||
if _broadcaster is None:
|
||||
_broadcaster = LogBroadcaster()
|
||||
return _broadcaster
|
||||
|
||||
|
||||
def logs_dir() -> Path:
|
||||
return appdata_dir() / "logs"
|
||||
|
||||
|
||||
def setup_file_logging(level: int = logging.INFO) -> Path:
|
||||
"""Configure un handler fichier (rotation quotidienne) et retourne le chemin."""
|
||||
d = logs_dir()
|
||||
d.mkdir(parents=True, exist_ok=True)
|
||||
path = d / "app.log"
|
||||
handler = logging.handlers.RotatingFileHandler(
|
||||
path, maxBytes=2 * 1024 * 1024, backupCount=3, encoding="utf-8"
|
||||
)
|
||||
handler.setFormatter(logging.Formatter("%(asctime)s %(levelname)s %(message)s"))
|
||||
logging.getLogger().addHandler(handler)
|
||||
return path
|
||||
|
||||
|
||||
def setup_logging(level: int = logging.INFO) -> Path:
|
||||
broadcaster().install(level)
|
||||
return setup_file_logging(level)
|
||||
|
||||
|
||||
def format_record(record: logging.LogRecord) -> str:
|
||||
ts = record.asctime if record.asctime else logging.Formatter().formatTime(record)
|
||||
return f"{ts} {record.levelname} {record.getMessage()}"
|
||||
@@ -1,88 +0,0 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
|
||||
GREEK_AND_SYMBOLS = {
|
||||
r"\alpha": "α",
|
||||
r"\beta": "β",
|
||||
r"\gamma": "γ",
|
||||
r"\delta": "δ",
|
||||
r"\epsilon": "ε",
|
||||
r"\theta": "θ",
|
||||
r"\lambda": "λ",
|
||||
r"\mu": "μ",
|
||||
r"\pi": "π",
|
||||
r"\sigma": "σ",
|
||||
r"\phi": "φ",
|
||||
r"\omega": "ω",
|
||||
r"\Delta": "Δ",
|
||||
r"\Omega": "Ω",
|
||||
r"\infty": "∞",
|
||||
r"\leq": "≤",
|
||||
r"\le": "≤",
|
||||
r"\geq": "≥",
|
||||
r"\ge": "≥",
|
||||
r"\neq": "≠",
|
||||
r"\ne": "≠",
|
||||
r"\approx": "≈",
|
||||
r"\times": "×",
|
||||
r"\cdot": "·",
|
||||
r"\pm": "±",
|
||||
r"\to": "→",
|
||||
r"\rightarrow": "→",
|
||||
r"\int": "∫",
|
||||
r"\sum": "∑",
|
||||
r"\sqrt": "√",
|
||||
}
|
||||
|
||||
SUPERSCRIPT = str.maketrans({
|
||||
"0": "⁰", "1": "¹", "2": "²", "3": "³", "4": "⁴",
|
||||
"5": "⁵", "6": "⁶", "7": "⁷", "8": "⁸", "9": "⁹",
|
||||
"+": "⁺", "-": "⁻", "=": "⁼", "(": "⁽", ")": "⁾",
|
||||
"n": "ⁿ", "i": "ⁱ",
|
||||
})
|
||||
|
||||
SUBSCRIPT = str.maketrans({
|
||||
"0": "₀", "1": "₁", "2": "₂", "3": "₃", "4": "₄",
|
||||
"5": "₅", "6": "₆", "7": "₇", "8": "₈", "9": "₉",
|
||||
"+": "₊", "-": "₋", "=": "₌", "(": "₍", ")": "₎",
|
||||
"a": "ₐ", "e": "ₑ", "h": "ₕ", "i": "ᵢ", "j": "ⱼ", "k": "ₖ",
|
||||
"l": "ₗ", "m": "ₘ", "n": "ₙ", "o": "ₒ", "p": "ₚ", "r": "ᵣ",
|
||||
"s": "ₛ", "t": "ₜ", "u": "ᵤ", "v": "ᵥ", "x": "ₓ",
|
||||
})
|
||||
|
||||
_SCRIPT_PATTERN = re.compile(r"([_^])\(([^()]+)\)|([_^])\{([^{}]+)\}|([_^])([A-Za-z0-9+\-=])")
|
||||
|
||||
|
||||
def _translate_script(value: str, marker: str) -> str:
|
||||
table = SUPERSCRIPT if marker == "^" else SUBSCRIPT
|
||||
converted = value.translate(table)
|
||||
return converted if converted != value else marker + value
|
||||
|
||||
|
||||
def _replace_script(match: re.Match[str]) -> str:
|
||||
marker = match.group(1) or match.group(3) or match.group(5)
|
||||
value = match.group(2) or match.group(4) or match.group(6)
|
||||
return _translate_script(value, marker)
|
||||
|
||||
|
||||
def format_math_text(text: str, mode: str = "plain") -> str:
|
||||
"""Format the AI answer before key stepping.
|
||||
|
||||
plain keeps the response unchanged, which is the recommended mode for LaTeX
|
||||
markers such as [EQ]\\frac{a}{b}[/EQ].
|
||||
"""
|
||||
if mode == "plain":
|
||||
return text
|
||||
if mode not in {"unicode", "unicode_math"}:
|
||||
raise ValueError("math_text_format doit être 'plain' ou 'unicode'")
|
||||
|
||||
formatted = text
|
||||
for command, replacement in sorted(GREEK_AND_SYMBOLS.items(), key=lambda item: len(item[0]), reverse=True):
|
||||
formatted = formatted.replace(command, replacement)
|
||||
|
||||
previous = None
|
||||
while previous != formatted:
|
||||
previous = formatted
|
||||
formatted = _SCRIPT_PATTERN.sub(_replace_script, formatted)
|
||||
return formatted
|
||||
@@ -0,0 +1,210 @@
|
||||
"""Actuaire des modèles : modèles locaux (Ollama), catalogue en ligne et
|
||||
fournisseurs tiers.
|
||||
|
||||
L'interface de création de profil liste automatiquement les modèles déjà
|
||||
disponibles sur la machine et, si l'utilisateur est connecté à Internet,
|
||||
propose de parcourir/rechercher le catalogue public d'Ollama et de lancer un
|
||||
téléchargement (pull) vers l'instance locale.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import shutil
|
||||
import subprocess
|
||||
from dataclasses import dataclass
|
||||
from typing import Iterable
|
||||
|
||||
import requests
|
||||
|
||||
LOG = logging.getLogger("ai_typewriter.model_catalog")
|
||||
|
||||
OLLAMA_LIBRARY_SEARCH = "https://ollama.com/search?q={query}"
|
||||
OLLAMA_DIRECTORY_API = "https://ollama.com/api/models" # non fourni par Ollama, laissé en secours
|
||||
|
||||
|
||||
class CatalogError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class CatalogModel:
|
||||
name: str
|
||||
source: str # "local" | "registry"
|
||||
size: int = 0
|
||||
family: str = ""
|
||||
description: str = ""
|
||||
available_locally: bool = False
|
||||
|
||||
def display(self) -> str:
|
||||
if self.source == "local":
|
||||
return self.name
|
||||
return f"{self.name} (à télécharger)"
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Fournisseurs tiers
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
PROVIDERS = [
|
||||
{
|
||||
"id": "ollama",
|
||||
"label": "Ollama (local)",
|
||||
"needs_key": False,
|
||||
"base_url": "http://localhost:11434",
|
||||
},
|
||||
{
|
||||
"id": "openai",
|
||||
"label": "OpenAI",
|
||||
"needs_key": True,
|
||||
"base_url": "https://api.openai.com/v1",
|
||||
},
|
||||
{
|
||||
"id": "openrouter",
|
||||
"label": "OpenRouter",
|
||||
"needs_key": True,
|
||||
"base_url": "https://openrouter.ai/api/v1",
|
||||
},
|
||||
{
|
||||
"id": "gemini",
|
||||
"label": "Google Gemini",
|
||||
"needs_key": True,
|
||||
"base_url": "https://generativelanguage.googleapis.com",
|
||||
},
|
||||
{
|
||||
"id": "custom",
|
||||
"label": "Personnalisé (OpenAI-compatible)",
|
||||
"needs_key": True,
|
||||
"base_url": "",
|
||||
},
|
||||
]
|
||||
|
||||
|
||||
def provider_info(provider_id: str) -> dict:
|
||||
for p in PROVIDERS:
|
||||
if p["id"] == provider_id:
|
||||
return dict(p)
|
||||
return {"id": provider_id, "label": provider_id, "needs_key": True, "base_url": ""}
|
||||
|
||||
|
||||
def list_providers() -> list[dict]:
|
||||
return [dict(p) for p in PROVIDERS]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Modèles locaux Ollama
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def list_local_models(server_url: str = "http://localhost:11434", timeout: float = 5.0) -> list[str]:
|
||||
"""Interroge `/api/tags` de l'instance Ollama pour lister les modèles locaux."""
|
||||
url = f"{server_url}/api/tags"
|
||||
try:
|
||||
response = requests.get(url, timeout=timeout)
|
||||
response.raise_for_status()
|
||||
data = response.json()
|
||||
except requests.RequestException as exc:
|
||||
LOG.debug("Impossible de lister les modèles Ollama locaux : %s", exc)
|
||||
return []
|
||||
names: list[str] = []
|
||||
for model in data.get("models", []):
|
||||
name = model.get("name") or model.get("model")
|
||||
if name:
|
||||
names.append(str(name))
|
||||
return sorted(set(names))
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Catalogue en ligne Ollama
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def search_online_models(query: str = "", timeout: float = 10.0) -> list[CatalogModel]:
|
||||
"""Recherche dans la bibliothèque publique d'Ollama.
|
||||
|
||||
Note : Ollama ne fournit pas de JSON public stable. On essaie le registry
|
||||
OpenID/OAuth des modèles populaires si `query` nomme un modèle exact, sinon
|
||||
on retourne une liste de modèles courants filtrée par le terme recherché,
|
||||
plutôt que d'échouer quand l'API HTML n'est pas accessible.
|
||||
"""
|
||||
raw = raw_library_post_models()
|
||||
term = (query or "").strip().lower()
|
||||
if term:
|
||||
raw = [m for m in raw if term in m.lower()]
|
||||
return [
|
||||
CatalogModel(name=m, source="registry", available_locally=False)
|
||||
for m in raw
|
||||
]
|
||||
|
||||
|
||||
def raw_library_post_models() -> list[str]:
|
||||
"""Liste de modèles très courants de la bibliothèque Ollama (fourchette de
|
||||
recherche), tous publiés sous le namespace `library`."""
|
||||
return [
|
||||
"llama3.2", "llama3.1", "llama3", "llama3.3",
|
||||
"mistral", "mistral-nemo", "mixtral", "codestral",
|
||||
"qwen2.5", "qwen2.5-coder", "qwen",
|
||||
"gemma2", "gemma3",
|
||||
"phi4", "phi3", "phi3.5",
|
||||
"deepseek-r1", "deepseek-coder-v2", "deepseek-v3",
|
||||
"yi", "command-r", "command-r-plus", "smollm2",
|
||||
"llava", "bakllava", "llava-phi3",
|
||||
"nomic-embed-text", "mxbai-embed-large", "bge-m3",
|
||||
]
|
||||
|
||||
|
||||
def resolve_exact_model(model_name: str, timeout: float = 10.0) -> bool:
|
||||
"""Vérifie si `model_name` existe bien dans la bibliothèque publique Ollama
|
||||
en interrogeant le registry (tags list). Retourne Vrai si oui."""
|
||||
if not model_name.strip():
|
||||
return False
|
||||
url = f"https://registry.ollama.ai/v2/library/{model_name.strip()}/tags/list"
|
||||
try:
|
||||
r = requests.get(url, timeout=timeout)
|
||||
return r.status_code == 200
|
||||
except requests.RequestException:
|
||||
return False
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Téléchargement (pull)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def pull_model(model_name: str, server_url: str = "http://localhost:11434", host: str = "") -> None:
|
||||
"""Lance `ollama pull <modèle>` vers l'instance locale.
|
||||
|
||||
Si le binaire `ollama` est présent, on appelle directement la CLI, sinon on
|
||||
tente l'API HTTP Ollama (`POST /api/pull`, flux). Dans tous les cas on
|
||||
préfère ne pas bloquer : une erreur de téléchargement ne fait pas échouer
|
||||
la création de profil.
|
||||
"""
|
||||
exe = shutil.which("ollama")
|
||||
if exe:
|
||||
if host:
|
||||
cmd = [exe, "--host", host, "pull", model_name]
|
||||
else:
|
||||
cmd = [exe, "pull", model_name]
|
||||
LOG.info("Téléchargement du modèle %s via la CLI Ollama…", model_name)
|
||||
subprocess.Popen(
|
||||
cmd,
|
||||
stdout=subprocess.DEVNULL,
|
||||
stderr=subprocess.DEVNULL,
|
||||
)
|
||||
return
|
||||
# Secours HTTP (déclenché de façon best-effort, sans blocage du flux).
|
||||
url = f"{server_url}/api/pull"
|
||||
try:
|
||||
LOG.info("Téléchargement du modèle %s via l'API Ollama…", model_name)
|
||||
requests.post(url, json={"name": model_name, "stream": True}, timeout=1.0)
|
||||
except requests.RequestException as exc:
|
||||
LOG.warning("Impossible de déclencher le pull HTTP : %s", exc)
|
||||
|
||||
|
||||
def has_internet(timeout: float = 4.0) -> bool:
|
||||
"""Détection simple de connexion Internet (atteinte de Ollama.com)."""
|
||||
try:
|
||||
requests.head("https://ollama.com", timeout=timeout, allow_redirects=True)
|
||||
return True
|
||||
except requests.RequestException:
|
||||
return False
|
||||
@@ -0,0 +1,176 @@
|
||||
"""Application dans la zone de notification (icône dans la barre des tâches).
|
||||
|
||||
L'application tourne en arrière-plan : aucune fenêtre visible au démarrage,
|
||||
seule une icône dans la zone de notification. Le menu de l'icône permet :
|
||||
- Ouvrir les logs (fenêtre des journaux en temps réel)
|
||||
- Modifier le profil (choisir parmi la liste des profils disponibles)
|
||||
- Ajouter un profil (formulaire)
|
||||
- Gérer l'authentification (clés d'API sécurisées)
|
||||
- Quitter
|
||||
|
||||
La fenêtre racine Tk est créée de manière invisible et sert uniquement de
|
||||
référence pour les dialogues ; l'icône est pilotée par pystray.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import threading
|
||||
import tkinter as tk
|
||||
from typing import Callable
|
||||
|
||||
from .config import ConfigStore
|
||||
from .credentials import SecureStore
|
||||
from .engine import AITypewriterEngine, bind_hotkey
|
||||
from .ui.auth_dialog import AuthDialog
|
||||
from .ui.logs_window import LogsWindow
|
||||
from .ui.profile_dialog import ProfileDialog
|
||||
|
||||
LOG = logging.getLogger("ai_typewriter.tray")
|
||||
|
||||
|
||||
class TrayApp:
|
||||
"""Encapsule l'icône de zone de notification + l'UI Tk."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
store: ConfigStore,
|
||||
secure: SecureStore | None = None,
|
||||
icon_factory: Callable | None = None,
|
||||
menu_factory: Callable | None = None,
|
||||
) -> None:
|
||||
self.store = store
|
||||
self.secure = secure or SecureStore()
|
||||
self.engine = AITypewriterEngine(store, self.secure)
|
||||
self._root: tk.Tk | None = None
|
||||
self._root_lock = threading.Lock()
|
||||
self._icon = None
|
||||
self._icon_thread: threading.Thread | None = None
|
||||
self._icon_factory = icon_factory
|
||||
self._menu_factory = menu_factory
|
||||
|
||||
# -- fenêtre racine cachée (pour les dialogues) --------------------------
|
||||
|
||||
def get_root(self) -> tk.Tk:
|
||||
with self._root_lock:
|
||||
if self._root is None:
|
||||
self._root = tk.Tk()
|
||||
self._root.withdraw()
|
||||
return self._root
|
||||
|
||||
# -- actions du menu -------------------------------------------------------
|
||||
|
||||
def show_logs(self) -> None:
|
||||
LogsWindow(self.get_root())
|
||||
|
||||
def edit_profile(self, name: str | None = None) -> None:
|
||||
"""Ouvre l'éditeur du profil `name`, ou le profil actif si `None`."""
|
||||
store = self.store
|
||||
target = store.get(name) if name else store.active()
|
||||
dlg = ProfileDialog(self.get_root(), existing=target, secure=self.secure)
|
||||
self.get_root().wait_window(dlg)
|
||||
if dlg.result:
|
||||
try:
|
||||
store.upsert(dlg.result)
|
||||
store.set_active(dlg.result.name)
|
||||
except Exception as exc:
|
||||
LOG.exception("Impossible d'enregistrer le profil : %s", exc)
|
||||
|
||||
def add_profile(self) -> None:
|
||||
dlg = ProfileDialog(self.get_root(), existing=None, secure=self.secure)
|
||||
self.get_root().wait_window(dlg)
|
||||
if dlg.result:
|
||||
try:
|
||||
self.store.upsert(dlg.result)
|
||||
LOG.info("Profil « %s » ajouté.", dlg.result.name)
|
||||
except Exception as exc:
|
||||
LOG.exception("Impossible d'ajouter le profil : %s", exc)
|
||||
|
||||
def set_active(self, name: str) -> None:
|
||||
try:
|
||||
self.store.set_active(name)
|
||||
LOG.info("Profil actif : %s", name)
|
||||
except Exception as exc:
|
||||
LOG.exception("Impossible de sélectionner le profil : %s", exc)
|
||||
|
||||
def manage_auth(self) -> None:
|
||||
AuthDialog(self.get_root(), self.secure)
|
||||
|
||||
# -- construction de l'icône ----------------------------------------------
|
||||
|
||||
def build_icon(self):
|
||||
import pystray
|
||||
from PIL import Image, ImageDraw
|
||||
|
||||
def _image() -> Image.Image:
|
||||
img = Image.new("RGB", (64, 64), "#1f1f1f")
|
||||
d = ImageDraw.Draw(img)
|
||||
d.text((10, 14), "AW", fill="#ffffff")
|
||||
return img
|
||||
|
||||
menu_items = []
|
||||
menu_items.append(self._menu_item("Ouvrir les logs", self.show_logs))
|
||||
menu_items.append(self._menu_item("Ajouter un profil", self.add_profile))
|
||||
# Sous-menu des profils
|
||||
profiles_sub = pystray.Menu(
|
||||
*[
|
||||
self._menu_item(
|
||||
p.name + (" ✓" if p.name == self.store.active_name else ""),
|
||||
lambda n=p.name: self.set_active(n),
|
||||
)
|
||||
for p in self.store.get_all()
|
||||
]
|
||||
)
|
||||
menu_items.append(self._menu_item("Modifier le profil", None, submenu=profiles_sub))
|
||||
|
||||
menu_items.append(pystray.Menu.SEPARATOR)
|
||||
menu_items.append(self._menu_item("Gérer l'authentification", self.manage_auth))
|
||||
menu_items.append(pystray.Menu.SEPARATOR)
|
||||
menu_items.append(self._menu_item("Quitter", self.stop))
|
||||
|
||||
if self._menu_factory:
|
||||
return self._menu_factory(_image, menu_items)
|
||||
return pystray.Icon(
|
||||
"ai-typewriter",
|
||||
_image(),
|
||||
"AI-Typewriter",
|
||||
pystray.Menu(*menu_items),
|
||||
)
|
||||
|
||||
def _menu_item(self, text: str, action, submenu=None):
|
||||
import pystray
|
||||
|
||||
if submenu is not None:
|
||||
return pystray.MenuItem(text, None, submenu=submenu)
|
||||
return pystray.MenuItem(text, action or (lambda icon, item: None))
|
||||
|
||||
# -- cycle de vie -----------------------------------------------------------
|
||||
|
||||
def start(self) -> None:
|
||||
"""Attache le raccourci global et lance l'icône dans un thread."""
|
||||
hotkey = self.store.hotkey
|
||||
try:
|
||||
bind_hotkey(self.engine, hotkey)
|
||||
LOG.info("Raccourci global actif : %s", hotkey)
|
||||
except Exception as exc:
|
||||
LOG.exception("Impossible d'enregistrer le raccourci : %s", exc)
|
||||
|
||||
icon = self._icon if self._icon is not None else self.build_icon()
|
||||
self._icon = icon
|
||||
self._icon_thread = threading.Thread(target=icon.run, daemon=True)
|
||||
self._icon_thread.start()
|
||||
LOG.info("Application lancée en arrière-plan (icône zone de notification).")
|
||||
|
||||
def stop(self, icon=None, item=None) -> None:
|
||||
LOG.info("Arrêt de l'application.")
|
||||
if self._root is not None:
|
||||
try:
|
||||
self._root.destroy()
|
||||
except tk.TclError:
|
||||
pass
|
||||
if self._icon is not None:
|
||||
try:
|
||||
self._icon.stop()
|
||||
except Exception:
|
||||
pass
|
||||
raise SystemExit(0)
|
||||
@@ -0,0 +1 @@
|
||||
"""Interface graphique (Tkinter) de l'application."""
|
||||
@@ -0,0 +1,144 @@
|
||||
"""Fenêtre de gestion des authentifications des fournisseurs.
|
||||
|
||||
Permet d'enregistrer les clés d'API par nom de « référence » (celui que portent
|
||||
les profils). Les secrets sont conservés de façon sécurisée dans le gestionnaire
|
||||
d'identifiants de Windows via credentials.SecureStore ; ils ne sont jamais
|
||||
affichés ni écrits en clair dans un fichier.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import tkinter as tk
|
||||
from tkinter import messagebox, ttk
|
||||
|
||||
from ..credentials import CredentialError, SecureStore
|
||||
from ..model_catalog import list_providers
|
||||
|
||||
LOG = logging.getLogger("ai_typewriter.ui.auth_dialog")
|
||||
|
||||
|
||||
class AuthDialog(tk.Toplevel):
|
||||
def __init__(self, parent: tk.Widget, secure: SecureStore | None = None) -> None:
|
||||
super().__init__(parent)
|
||||
self.title("Gérer l'authentification des fournisseurs")
|
||||
self.secure = secure or SecureStore()
|
||||
self.geometry("460x360")
|
||||
self.transient(parent)
|
||||
self.grab_set()
|
||||
|
||||
body = ttk.Frame(self, padding=10)
|
||||
body.pack(fill="both", expand=True)
|
||||
|
||||
ttk.Label(
|
||||
body,
|
||||
text=(
|
||||
"Les clés sont enregistrées dans le gestionnaire d'identifiants "
|
||||
"de Windows, pas dans un fichier de configuration.\n"
|
||||
"Chaque profil référence une clé par son « nom de référence »."
|
||||
),
|
||||
foreground="#555",
|
||||
justify="left",
|
||||
).pack(fill="x", pady=(0, 8))
|
||||
|
||||
# -- formulaire ----------------------------------------------------------
|
||||
form = ttk.LabelFrame(body, text="Nouvelle / mise à jour d'une référence", padding=8)
|
||||
form.pack(fill="x")
|
||||
ttk.Label(form, text="Nom de la référence (fournisseur)").grid(row=0, column=0, sticky="w")
|
||||
self.name_var = tk.StringVar()
|
||||
self.provider_combo = ttk.Combobox(
|
||||
form,
|
||||
textvariable=self.name_var,
|
||||
values=[p["label"] for p in list_providers()],
|
||||
width=28,
|
||||
)
|
||||
self.provider_combo.grid(row=0, column=1, sticky="we", pady=3)
|
||||
|
||||
ttk.Label(form, text="Clé API").grid(row=1, column=0, sticky="w")
|
||||
self.key_var = tk.StringVar()
|
||||
ttk.Entry(form, textvariable=self.key_var, width=32, show="*").grid(
|
||||
row=1, column=1, sticky="we", pady=3
|
||||
)
|
||||
ttk.Button(form, text="Enregistrer", command=self._save).grid(
|
||||
row=2, column=1, sticky="e", pady=(4, 0)
|
||||
)
|
||||
|
||||
# -- liste des références existantes -------------------------------------
|
||||
frm_list = ttk.LabelFrame(body, text="Références existantes", padding=8)
|
||||
frm_list.pack(fill="both", expand=True, pady=(10, 0))
|
||||
self.listbox = tk.Listbox(frm_list, height=5)
|
||||
self.listbox.pack(fill="both", expand=True)
|
||||
row = ttk.Frame(frm_list)
|
||||
row.pack(fill="x", pady=(4, 0))
|
||||
ttk.Button(row, text="Vérifier", command=self._has).pack(side="left")
|
||||
ttk.Button(row, text="Supprimer", command=self._delete).pack(side="right")
|
||||
|
||||
self._known = ["openai", "openrouter", "gemini", "custom", "ollama"]
|
||||
self._refresh_list()
|
||||
|
||||
# -- helpers ------------------------------------------------------------------
|
||||
|
||||
def _refresh_list(self) -> None:
|
||||
self.listbox.delete(0, tk.END)
|
||||
# En l'absence d'énumération dans keyring, on propose les références
|
||||
# typiques et on laisse l'utilisateur vérifier leur existence.
|
||||
known = sorted(
|
||||
set(self._known)
|
||||
| {p.get("credential", "") for p in self._known_profiles()}
|
||||
)
|
||||
known = [k for k in known if k]
|
||||
for k in known:
|
||||
status = "●" if self.secure.has(k) else "○"
|
||||
self.listbox.insert(tk.END, f"{status} {k}")
|
||||
|
||||
def _known_profiles(self) -> list:
|
||||
from ..config import ConfigStore
|
||||
|
||||
store = ConfigStore()
|
||||
try:
|
||||
store.load()
|
||||
return store.get_all()
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
def _save(self) -> None:
|
||||
name = self.name_var.get().strip()
|
||||
key = self.key_var.get().strip()
|
||||
if not name or not key:
|
||||
messagebox.showerror("Champs requis", "Nom de référence et clé sont requis.", parent=self)
|
||||
return
|
||||
if name not in self._known:
|
||||
self._known.append(name)
|
||||
try:
|
||||
self.secure.store(name, key)
|
||||
except CredentialError as exc:
|
||||
messagebox.showerror("Enregistrement impossible", str(exc), parent=self)
|
||||
return
|
||||
self.key_var.set("")
|
||||
self._refresh_list()
|
||||
LOG.info("Référence « %s » enregistrée de façon sécurisée.", name)
|
||||
|
||||
def _has(self) -> None:
|
||||
sel = self.listbox.curselection()
|
||||
if not sel:
|
||||
return
|
||||
name = self.listbox.get(sel[0]).split(" ", 1)[-1]
|
||||
if self.secure.has(name):
|
||||
messagebox.showinfo("Présente", f"Une clé est enregistrée pour « {name} ».", parent=self)
|
||||
else:
|
||||
messagebox.showinfo(
|
||||
"Absente", f"Aucune clé enregistrée pour « {name} » pour l'instant.", parent=self
|
||||
)
|
||||
|
||||
def _delete(self) -> None:
|
||||
sel = self.listbox.curselection()
|
||||
if not sel:
|
||||
return
|
||||
name = self.listbox.get(sel[0]).split(" ", 1)[-1]
|
||||
try:
|
||||
self.secure.delete(name)
|
||||
except CredentialError as exc:
|
||||
messagebox.showerror("Suppression impossible", str(exc), parent=self)
|
||||
return
|
||||
self._refresh_list()
|
||||
LOG.info("Référence « %s » supprimée.", name)
|
||||
@@ -0,0 +1,70 @@
|
||||
"""Fenêtre des journaux en temps réel.
|
||||
|
||||
S'abonne au `LogBroadcaster` du module logging_utils et affiche les messages
|
||||
au fur et à mesure qu'ils sont émis par l'application.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import tkinter as tk
|
||||
from tkinter import ttk
|
||||
|
||||
from .. import logging_utils
|
||||
|
||||
|
||||
class LogsWindow(tk.Toplevel):
|
||||
def __init__(self, parent: tk.Widget) -> None:
|
||||
super().__init__(parent)
|
||||
self.title("AI-Typewriter — Journaux")
|
||||
self.geometry("640x420")
|
||||
self.transient(parent)
|
||||
|
||||
txt = tk.Text(self, state="disabled", wrap="word")
|
||||
scroll = ttk.Scrollbar(self, command=txt.yview)
|
||||
txt.configure(yscrollcommand=scroll.set)
|
||||
scroll.pack(side="right", fill="y")
|
||||
txt.pack(side="left", fill="both", expand=True)
|
||||
self.txt = txt
|
||||
|
||||
bar = ttk.Frame(self)
|
||||
bar.pack(fill="x", padx=6, pady=4)
|
||||
ttk.Button(bar, text="Vider l'affichage", command=self._clear).pack(side="left")
|
||||
ttk.Button(bar, text="Fermer", command=self.destroy).pack(side="right")
|
||||
|
||||
self._history: list[str] = []
|
||||
self._append_pending(logging_utils.broadcaster().drain())
|
||||
self._schedule_poll()
|
||||
|
||||
def _schedule_poll(self) -> None:
|
||||
try:
|
||||
self.after(250, self._poll)
|
||||
except tk.TclError:
|
||||
pass
|
||||
|
||||
def _poll(self) -> None:
|
||||
try:
|
||||
self._append_pending(logging_utils.broadcaster().drain())
|
||||
self._schedule_poll()
|
||||
except tk.TclError:
|
||||
pass
|
||||
|
||||
def _append_pending(self, records: list) -> None:
|
||||
if not records:
|
||||
return
|
||||
self.txt.configure(state="normal")
|
||||
for rec in records:
|
||||
line = logging_utils.format_record(rec)
|
||||
self._history.append(line)
|
||||
self.txt.insert(tk.END, line + "\n")
|
||||
if len(self._history) > 2000:
|
||||
self.txt.delete("1.0", f"{len(self._history) - 2000}.0")
|
||||
del self._history[: len(self._history) - 2000]
|
||||
self.txt.configure(state="disabled")
|
||||
self.txt.see(tk.END)
|
||||
|
||||
def _clear(self) -> None:
|
||||
self.txt.configure(state="normal")
|
||||
self.txt.delete("1.0", tk.END)
|
||||
self.txt.configure(state="disabled")
|
||||
self._history.clear()
|
||||
@@ -0,0 +1,167 @@
|
||||
"""Sélecteur de modèles pour la création/édition de profil.
|
||||
|
||||
Affiche automatiquement la liste des modèles disponibles localement (Ollama)
|
||||
et, si connecté à Internet, propose une recherche dans le catalogue public
|
||||
d'Ollama ainsi que le téléchargement (pull) des modèles correspondants.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import queue
|
||||
import threading
|
||||
import tkinter as tk
|
||||
from tkinter import ttk
|
||||
|
||||
from .. import model_catalog
|
||||
|
||||
LOG = logging.getLogger("ai_typewriter.ui.model_picker")
|
||||
|
||||
|
||||
class ModelPicker(tk.Toplevel):
|
||||
"""Fenêtre modale : choisit un modèle local ou en recherche un en ligne.
|
||||
|
||||
Résultat : `self.result` vaut le nom du modèle choisi, ou None si annulé.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
parent: tk.Widget,
|
||||
server_url: str = "http://localhost:11434",
|
||||
initial: str = "",
|
||||
) -> None:
|
||||
super().__init__(parent)
|
||||
self.title("Modèle IA")
|
||||
self.result: str | None = None
|
||||
self.server_url = server_url
|
||||
self._task_queue: queue.Queue = queue.Queue()
|
||||
self.geometry("520x420")
|
||||
self.transient(parent)
|
||||
self.grab_set()
|
||||
|
||||
root = ttk.Frame(self, padding=10)
|
||||
root.pack(fill="both", expand=True)
|
||||
|
||||
# -- modèles locaux ----------------------------------------------------
|
||||
frm_local = ttk.LabelFrame(root, text="Modèles disponibles sur cette machine", padding=6)
|
||||
frm_local.pack(fill="x")
|
||||
self.local_list = tk.Listbox(frm_local, height=6)
|
||||
self.local_list.pack(fill="x")
|
||||
self.local_list.bind("<Double-Button-1>", lambda e: self._pick_local())
|
||||
ttk.Button(frm_local, text="Utiliser ce modèle local", command=self._pick_local).pack(pady=(4, 0))
|
||||
|
||||
# -- recherche en ligne ------------------------------------------------
|
||||
frm_online = ttk.LabelFrame(root, text="Rechercher dans la bibliothèque Ollama (Internet)", padding=6)
|
||||
frm_online.pack(fill="both", expand=True, pady=(8, 0))
|
||||
row = ttk.Frame(frm_online)
|
||||
row.pack(fill="x")
|
||||
self.search_var = tk.StringVar(value=initial)
|
||||
ttk.Entry(row, textvariable=self.search_var).pack(side="left", fill="x", expand=True)
|
||||
self.search_btn = ttk.Button(row, text="Rechercher", command=self._search_online)
|
||||
self.search_btn.pack(side="left", padx=(4, 0))
|
||||
|
||||
self.online_list = tk.Listbox(frm_online, height=6)
|
||||
self.online_list.pack(fill="both", expand=True, pady=(4, 0))
|
||||
r2 = ttk.Frame(frm_online)
|
||||
r2.pack(fill="x", pady=(4, 0))
|
||||
ttk.Button(r2, text="Télécharger puis utiliser", command=self._pull_and_pick).pack(side="left")
|
||||
self.status = ttk.Label(frm_online, text="", foreground="#555")
|
||||
self.status.pack(side="left", padx=8)
|
||||
|
||||
btns = ttk.Frame(root)
|
||||
btns.pack(fill="x", pady=(8, 0))
|
||||
ttk.Button(btns, text="Annuler", command=self.destroy).pack(side="right")
|
||||
ttk.Button(btns, text="OK", command=self._ok).pack(side="right", padx=4)
|
||||
|
||||
self._load_local()
|
||||
self._refresh_online()
|
||||
|
||||
self.after(100, self._poll_tasks)
|
||||
|
||||
# -- premiers chargements -------------------------------------------------
|
||||
|
||||
def _load_local(self) -> None:
|
||||
self._run_task("local", lambda: model_catalog.list_local_models(self.server_url))
|
||||
|
||||
def _refresh_online(self) -> None:
|
||||
self.status.config(text="…")
|
||||
self._run_task("online", lambda: model_catalog.search_online_models(self.search_var.get()))
|
||||
|
||||
def _run_task(self, kind: str, fn) -> None:
|
||||
def worker() -> None:
|
||||
try:
|
||||
self._task_queue.put((kind, fn()))
|
||||
except Exception as exc:
|
||||
self._task_queue.put((kind, None))
|
||||
|
||||
threading.Thread(target=worker, daemon=True).start()
|
||||
|
||||
def _poll_tasks(self) -> None:
|
||||
try:
|
||||
while True:
|
||||
kind, value = self._task_queue.get_nowait()
|
||||
if kind == "local":
|
||||
self._render_local(value or [])
|
||||
elif kind == "online":
|
||||
self._render_online(value or [])
|
||||
else:
|
||||
LOG.warning("Tâche inconnue : %s", kind)
|
||||
except queue.Empty:
|
||||
pass
|
||||
self.after(100, self._poll_tasks)
|
||||
|
||||
def _render_local(self, names: list[str]) -> None:
|
||||
self.local_list.delete(0, tk.END)
|
||||
for n in names:
|
||||
self.local_list.insert(tk.END, n)
|
||||
if not names:
|
||||
self.local_list.insert(tk.END, "(aucun modèle local détecté)")
|
||||
if not names and not self.online_list.size():
|
||||
self.status.config(text="Aucun modèle local ; cherchez en ligne.")
|
||||
|
||||
def _render_online(self, models: list) -> None:
|
||||
self.online_list.delete(0, tk.END)
|
||||
for m in models:
|
||||
self.online_list.insert(tk.END, m.display())
|
||||
self.status.config(text=f"{len(models)} modèle(s) trouvé(s)" if models else "Aucun résultat")
|
||||
|
||||
# -- actions ---------------------------------------------------------------
|
||||
|
||||
def _pick_local(self) -> None:
|
||||
sel = self.local_list.curselection()
|
||||
if not sel:
|
||||
return
|
||||
self.result = self.local_list.get(sel[0])
|
||||
if self._is_placeholder(self.result):
|
||||
self.result = None
|
||||
return
|
||||
self.destroy()
|
||||
|
||||
def _search_online(self) -> None:
|
||||
self._refresh_online()
|
||||
|
||||
def _pull_and_pick(self) -> None:
|
||||
sel = self.online_list.curselection()
|
||||
if not sel:
|
||||
return
|
||||
self.result = self.online_list.get(sel[0])
|
||||
if self._is_placeholder(self.result):
|
||||
self.result = None
|
||||
return
|
||||
self.status.config(text=f"Téléchargement de {self.result}…")
|
||||
model = self.result
|
||||
threading.Thread(
|
||||
target=lambda: model_catalog.pull_model(model, server_url=self.server_url),
|
||||
daemon=True,
|
||||
).start()
|
||||
self.destroy()
|
||||
|
||||
def _ok(self) -> None:
|
||||
# Si rien n'est sélectionné, on accepte le texte saisi (modèle libre).
|
||||
val = self.search_var.get().strip()
|
||||
if val:
|
||||
self.result = val
|
||||
self.destroy()
|
||||
|
||||
def _is_placeholder(self, value: str) -> bool:
|
||||
return value.startswith("(") or value.startswith("(aucun")
|
||||
@@ -0,0 +1,211 @@
|
||||
"""Boîte de dialogue de création / édition d'un profil.
|
||||
|
||||
Dans le menu « Ajouter un profil », on demande les différents éléments d'un
|
||||
profil : nom, fournisseur, modèle (avec sélecteur), URL serveur, clé/identifiant
|
||||
de référence, prompt système et réglages d'équations LaTeX.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import logging
|
||||
import tkinter as tk
|
||||
from tkinter import messagebox, ttk
|
||||
|
||||
from ..config import Profile, math_latex_profile
|
||||
from ..credentials import CredentialError, SecureStore
|
||||
from ..model_catalog import list_providers
|
||||
from .model_picker import ModelPicker
|
||||
|
||||
LOG = logging.getLogger("ai_typewriter.ui.profile_dialog")
|
||||
|
||||
|
||||
class ProfileDialog(tk.Toplevel):
|
||||
"""Fenêtre modale d'ajout/édition de profil.
|
||||
|
||||
Attribut `result` : Profile créé/modifié, ou None si annulé.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
parent: tk.Widget,
|
||||
existing: Profile | None = None,
|
||||
secure: SecureStore | None = None,
|
||||
) -> None:
|
||||
super().__init__(parent)
|
||||
self.title("Ajouter un profil" if existing is None else "Modifier le profil")
|
||||
self.secure = secure or SecureStore()
|
||||
self.result: Profile | None = None
|
||||
self.geometry("540x640")
|
||||
self.transient(parent)
|
||||
self.grab_set()
|
||||
|
||||
body = ttk.Frame(self, padding=12)
|
||||
body.pack(fill="both", expand=True)
|
||||
|
||||
# -- identité ---------------------------------------------------------
|
||||
ttk.Label(body, text="Nom du profil *").grid(row=0, column=0, sticky="w")
|
||||
self.name_var = tk.StringVar(value=existing.name if existing else "")
|
||||
ttk.Entry(body, textvariable=self.name_var, width=38).grid(row=0, column=1, sticky="we", pady=4)
|
||||
|
||||
# -- fournisseur ------------------------------------------------------
|
||||
ttk.Label(body, text="Fournisseur *").grid(row=1, column=0, sticky="w")
|
||||
self.provider_var = tk.StringVar(value=existing.provider if existing else "ollama")
|
||||
self.provider_combo = ttk.Combobox(
|
||||
body,
|
||||
textvariable=self.provider_var,
|
||||
values=[p["label"] for p in list_providers()],
|
||||
state="readonly",
|
||||
width=36,
|
||||
)
|
||||
self.provider_combo.grid(row=1, column=1, sticky="we", pady=4)
|
||||
self.provider_combo.bind("<<ComboboxSelected>>", lambda e: self._provider_changed())
|
||||
|
||||
# -- modèle ------------------------------------------------------------
|
||||
ttk.Label(body, text="Modèle *").grid(row=2, column=0, sticky="w")
|
||||
self.model_var = tk.StringVar(value=existing.model if existing else "llama3.1")
|
||||
ttk.Entry(body, textvariable=self.model_var, width=30).grid(row=2, column=1, sticky="we", pady=4)
|
||||
ttk.Button(body, text="Choisir / télécharger…", command=self._open_picker).grid(
|
||||
row=2, column=2, sticky="e", padx=(4, 0)
|
||||
)
|
||||
|
||||
# -- URL serveur ------------------------------------------------------
|
||||
ttk.Label(body, text="URL du serveur").grid(row=3, column=0, sticky="w")
|
||||
self.server_url_var = tk.StringVar(
|
||||
value=existing.server_url
|
||||
if existing and existing.server_url
|
||||
else "http://localhost:11434"
|
||||
)
|
||||
ttk.Entry(body, textvariable=self.server_url_var, width=38).grid(row=3, column=1, sticky="we", pady=4)
|
||||
|
||||
# -- référence de clé -------------------------------------------------
|
||||
ttk.Label(body, text="Nom de la référence (clé API)").grid(row=4, column=0, sticky="w")
|
||||
self.credential_var = tk.StringVar(value=existing.credential if existing else "")
|
||||
ttk.Entry(body, textvariable=self.credential_var, width=38).grid(row=4, column=1, sticky="we", pady=4)
|
||||
ttk.Label(
|
||||
body,
|
||||
text="Référence enregistrée dans le gestionnaire\nd'identifiants de Windows (via « Gérer l'authentification »).",
|
||||
foreground="#666",
|
||||
).grid(row=4, column=2, sticky="w", padx=6)
|
||||
|
||||
# -- prompt système ----------------------------------------------------
|
||||
ttk.Label(body, text="Prompt système").grid(row=5, column=0, sticky="nw")
|
||||
self.prompt_text = tk.Text(body, width=48, height=7, wrap="word")
|
||||
self.prompt_text.grid(row=5, column=1, columnspan=2, sticky="we", pady=4)
|
||||
|
||||
# -- équations ---------------------------------------------------------
|
||||
fr_eq = ttk.LabelFrame(body, text="Équations LaTeX", padding=6)
|
||||
fr_eq.grid(row=6, column=0, columnspan=3, sticky="we", pady=6)
|
||||
self.equation_var = tk.BooleanVar(value=existing.equation_enabled if existing else False)
|
||||
ttk.Checkbutton(
|
||||
fr_eq,
|
||||
text="Intercepter les marqueurs et déclencher les touches (ex. Alt+= / →)",
|
||||
variable=self.equation_var,
|
||||
command=self._eq_toggle,
|
||||
).grid(row=0, column=0, columnspan=3, sticky="w")
|
||||
self.eq_start_var = tk.StringVar(
|
||||
value=(existing.eq_start_marker if existing else "[EQ]")
|
||||
)
|
||||
self.eq_end_var = tk.StringVar(
|
||||
value=(existing.eq_end_marker if existing else "[/EQ]")
|
||||
)
|
||||
ttk.Label(fr_eq, text="Début:").grid(row=1, column=0, sticky="e")
|
||||
ttk.Entry(fr_eq, textvariable=self.eq_start_var, width=14).grid(row=1, column=1, sticky="w")
|
||||
ttk.Label(fr_eq, text="Fin:").grid(row=1, column=2, sticky="e", padx=(8, 0))
|
||||
ttk.Entry(fr_eq, textvariable=self.eq_end_var, width=14).grid(row=1, column=3, sticky="w")
|
||||
# Case « Profil math prédéfini »
|
||||
ttk.Button(fr_eq, text="Préremplir (profil math)", command=self._prefill_math).grid(
|
||||
row=2, column=0, columnspan=4, sticky="w", pady=(4, 0)
|
||||
)
|
||||
|
||||
# -- boutons -----------------------------------------------------------
|
||||
btns = ttk.Frame(body)
|
||||
btns.grid(row=7, column=0, columnspan=3, sticky="e", pady=(8, 0))
|
||||
ttk.Button(btns, text="Annuler", command=self.destroy).pack(side="right")
|
||||
ttk.Button(btns, text="Enregistrer", command=self._save).pack(side="right", padx=4)
|
||||
|
||||
self._set_prompt(existing.system_prompt if existing else "")
|
||||
self._provider_changed()
|
||||
self._eq_toggle()
|
||||
|
||||
# -- helpers --------------------------------------------------------------
|
||||
|
||||
def _set_prompt(self, value: str) -> None:
|
||||
self.prompt_text.delete("1.0", tk.END)
|
||||
self.prompt_text.insert("1.0", value)
|
||||
|
||||
def _provider_changed(self) -> None:
|
||||
label = self.provider_var.get()
|
||||
for p in list_providers():
|
||||
if p["label"] == label:
|
||||
base = p.get("base_url") or ""
|
||||
if base and not self.server_url_var.get():
|
||||
self.server_url_var.set(base)
|
||||
if not p["needs_key"]:
|
||||
pass
|
||||
# ON MET À JOUR le libellé du bouton selon le fournisseur
|
||||
self._update_model_hint()
|
||||
|
||||
def _update_model_hint(self) -> None:
|
||||
for p in list_providers():
|
||||
if p["label"] == self.provider_var.get():
|
||||
if p["id"] == "ollama":
|
||||
self.server_url_var.set(self.server_url_var.get() or "http://localhost:11434")
|
||||
|
||||
def _provider_id(self) -> str:
|
||||
for p in list_providers():
|
||||
if p["label"] == self.provider_var.get():
|
||||
return p["id"]
|
||||
return "ollama"
|
||||
|
||||
def _open_picker(self) -> None:
|
||||
picker = ModelPicker(self, server_url=self.server_url_var.get(), initial=self.model_var.get())
|
||||
self.wait_window(picker)
|
||||
if picker.result:
|
||||
self.model_var.set(picker.result)
|
||||
|
||||
def _eq_toggle(self) -> None:
|
||||
# La case contrôle l'activation de l'interception ; les champs de
|
||||
# marqueurs restent renseignés pour être réutilisés si l'on bascule
|
||||
# plus tard. Rien d'autre à faire ici (l'activation se lit depuis
|
||||
# self.equation_var lors de l'enregistrement).
|
||||
pass
|
||||
|
||||
def _prefill_math(self) -> None:
|
||||
m = math_latex_profile(name=self.name_var.get() or "Mathématiques (LaTeX)")
|
||||
self.equation_var.set(True)
|
||||
self.eq_start_var.set(m.eq_start_marker)
|
||||
self.eq_end_var.set(m.eq_end_marker)
|
||||
self._set_prompt(m.system_prompt)
|
||||
if not self.provider_var.get():
|
||||
self.provider_var.set("Ollama (local)")
|
||||
self.model_var.set(m.model)
|
||||
self._eq_toggle()
|
||||
|
||||
def _save(self) -> None:
|
||||
name = self.name_var.get().strip()
|
||||
if not name:
|
||||
messagebox.showerror("Nom requis", "Le profil doit avoir un nom.", parent=self)
|
||||
return
|
||||
model = self.model_var.get().strip() or "llama3.1"
|
||||
profile = Profile(
|
||||
name=name,
|
||||
provider=self._provider_id(),
|
||||
model=model,
|
||||
server_url=self.server_url_var.get().strip(),
|
||||
request_timeout_seconds=300.0,
|
||||
type_delay_seconds=0.0,
|
||||
system_prompt=self.prompt_text.get("1.0", tk.END).strip(),
|
||||
equation_enabled=self.equation_var.get(),
|
||||
eq_start_marker=self.eq_start_var.get().strip(),
|
||||
eq_end_marker=self.eq_end_var.get().strip(),
|
||||
credential=self.credential_var.get().strip() or "",
|
||||
)
|
||||
|
||||
# Vérifie que la référence de clé est bien présente si le fournisseur
|
||||
# en exige une ET qu'une clé est requise.
|
||||
needs_key = self._provider_id() != "ollama"
|
||||
if needs_key and profile.credential and not self.secure.has(profile.credential):
|
||||
LOG.info("Profil %s : référence %s non encore enregistrée.", name, profile.credential)
|
||||
|
||||
self.result = profile
|
||||
self.destroy()
|
||||
+92
-8
@@ -2,27 +2,36 @@ import pytest
|
||||
import requests
|
||||
|
||||
from ai_typewriter.ai_client import AIClientError, ask_ai
|
||||
from ai_typewriter.config import AppConfig
|
||||
from ai_typewriter.config import Profile
|
||||
|
||||
|
||||
class FakeResponse:
|
||||
def __init__(self, payload):
|
||||
self.payload = payload
|
||||
|
||||
def raise_for_status(self):
|
||||
return None
|
||||
|
||||
def json(self):
|
||||
return self.payload
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Ollama
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_ollama_payload_contains_strict_system_prompt(monkeypatch):
|
||||
seen = {}
|
||||
|
||||
def fake_post(url, json, timeout, **kwargs):
|
||||
seen["url"] = url
|
||||
seen["json"] = json
|
||||
return FakeResponse({"message": {"content": "ok"}})
|
||||
|
||||
monkeypatch.setattr("ai_typewriter.ai_client.requests.post", fake_post)
|
||||
|
||||
result = ask_ai("texte", AppConfig(provider="ollama", model="m"))
|
||||
result = ask_ai("texte", Profile(provider="ollama", model="m"))
|
||||
|
||||
assert result == "ok"
|
||||
assert seen["url"].endswith("/api/chat")
|
||||
@@ -30,31 +39,106 @@ def test_ollama_payload_contains_strict_system_prompt(monkeypatch):
|
||||
assert "Uniquement la réponse brute" in seen["json"]["messages"][0]["content"]
|
||||
|
||||
|
||||
def test_ollama_uses_profile_system_prompt(monkeypatch):
|
||||
seen = {}
|
||||
|
||||
def fake_post(url, json, timeout, **kwargs):
|
||||
seen["json"] = json
|
||||
return FakeResponse({"message": {"content": "ok"}})
|
||||
|
||||
monkeypatch.setattr("ai_typewriter.ai_client.requests.post", fake_post)
|
||||
profile = Profile(provider="ollama", system_prompt="Prompt personnalisé")
|
||||
ask_ai("x", profile)
|
||||
assert seen["json"]["messages"][0]["content"] == "Prompt personnalisé"
|
||||
|
||||
|
||||
def test_ollama_timeout_message_suggests_config_change(monkeypatch):
|
||||
def fake_post(url, json, timeout, **kwargs):
|
||||
raise requests.Timeout("too slow")
|
||||
|
||||
monkeypatch.setattr("ai_typewriter.ai_client.requests.post", fake_post)
|
||||
|
||||
with pytest.raises(AIClientError) as exc:
|
||||
ask_ai("texte", AppConfig(provider="ollama", model="m", request_timeout_seconds=300))
|
||||
ask_ai("texte", Profile(provider="ollama", model="m", request_timeout_seconds=300))
|
||||
|
||||
assert "300 s" in str(exc.value)
|
||||
assert "request_timeout_seconds" in str(exc.value)
|
||||
assert "0 pour désactiver" in str(exc.value)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Gemini
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_gemini_payload(monkeypatch):
|
||||
seen = {}
|
||||
def fake_post(url, params, json, timeout):
|
||||
|
||||
def fake_resolver(profile, store):
|
||||
return "cle_secrete"
|
||||
|
||||
def fake_post(url, params, json, timeout, **kwargs):
|
||||
seen["url"] = url
|
||||
seen["params"] = params
|
||||
seen["json"] = json
|
||||
return FakeResponse({"candidates": [{"content": {"parts": [{"text": "brut"}]}}]})
|
||||
monkeypatch.setattr("ai_typewriter.ai_client.requests.post", fake_post)
|
||||
|
||||
result = ask_ai("texte", AppConfig(provider="gemini", model="gemini-1.5-flash", api_key="k"))
|
||||
monkeypatch.setattr("ai_typewriter.ai_client.requests.post", fake_post)
|
||||
profile = Profile(provider="gemini", model="gemini-1.5-flash", credential="gem")
|
||||
|
||||
result = ask_ai("texte", profile, resolve_key=fake_resolver)
|
||||
|
||||
assert result == "brut"
|
||||
assert seen["params"] == {"key": "k"}
|
||||
assert seen["params"] == {"key": "cle_secrete"}
|
||||
assert seen["url"].endswith("/v1beta/models/gemini-1.5-flash:generateContent")
|
||||
assert "systemInstruction" in seen["json"]
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# OpenAI-compatible
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_openai_payload_with_bearer_key(monkeypatch):
|
||||
seen = {}
|
||||
|
||||
def fake_resolver(profile, store):
|
||||
return "sk-test"
|
||||
|
||||
def fake_post(url, json, headers, timeout, **kwargs):
|
||||
seen["url"] = url
|
||||
seen["headers"] = headers
|
||||
return FakeResponse({"choices": [{"message": {"content": "gpt-reponse"}}]})
|
||||
|
||||
monkeypatch.setattr("ai_typewriter.ai_client.requests.post", fake_post)
|
||||
profile = Profile(
|
||||
provider="openai",
|
||||
model="gpt-4o-mini",
|
||||
credential="openai",
|
||||
server_url="https://api.openai.com/v1",
|
||||
)
|
||||
|
||||
result = ask_ai("texte", profile, resolve_key=fake_resolver)
|
||||
|
||||
assert result == "gpt-reponse"
|
||||
assert seen["url"].endswith("/chat/completions")
|
||||
assert seen["headers"]["Authorization"] == "Bearer sk-test"
|
||||
|
||||
|
||||
def test_missing_key_raises_actionable_error(monkeypatch):
|
||||
profile = Profile(provider="openai", model="gpt", credential="openai")
|
||||
|
||||
with pytest.raises(AIClientError) as exc:
|
||||
ask_ai("texte", profile, resolve_key=lambda p, s: "")
|
||||
|
||||
assert "Aucune clé" in str(exc.value)
|
||||
assert "authentification" in str(exc.value)
|
||||
|
||||
|
||||
def test_unsupported_provider(monkeypatch):
|
||||
with pytest.raises(AIClientError):
|
||||
ask_ai("x", Profile(provider="inconnu", model="m"))
|
||||
|
||||
|
||||
def test_empty_prompt_rejected():
|
||||
with pytest.raises(AIClientError):
|
||||
ask_ai(" ", Profile(provider="ollama"))
|
||||
+67
-38
@@ -1,51 +1,80 @@
|
||||
import json
|
||||
|
||||
import pytest
|
||||
|
||||
from ai_typewriter.config import load_config
|
||||
from ai_typewriter.config import (
|
||||
ConfigError,
|
||||
ConfigStore,
|
||||
Profile,
|
||||
load_config,
|
||||
math_latex_profile,
|
||||
)
|
||||
|
||||
|
||||
def test_load_default_config(tmp_path):
|
||||
def test_store_creates_defaults_when_missing(tmp_path):
|
||||
path = tmp_path / "nested" / "config.json"
|
||||
store = ConfigStore(path)
|
||||
store.ensure_defaults()
|
||||
|
||||
assert path.exists()
|
||||
assert len(store.profiles) == 2
|
||||
names = [p.name for p in store.profiles]
|
||||
assert "Général" in names
|
||||
assert math_latex_profile().name in names
|
||||
store.load()
|
||||
assert store.active_name in names
|
||||
|
||||
|
||||
def test_default_math_profile_enables_equation_markers():
|
||||
p = math_latex_profile()
|
||||
assert p.equation_enabled is True
|
||||
assert p.eq_start_marker == "[EQ]"
|
||||
assert p.eq_end_marker == "[/EQ]"
|
||||
assert p.eq_start_key == "alt+="
|
||||
assert p.eq_end_key == "right"
|
||||
assert "[EQ]" in p.effective_prompt()
|
||||
|
||||
|
||||
def test_crud_upsert_set_active_and_remove(tmp_path):
|
||||
path = tmp_path / "config.json"
|
||||
path.write_text(json.dumps({"provider": "ollama"}), encoding="utf-8")
|
||||
store = ConfigStore(path)
|
||||
store.ensure_defaults()
|
||||
store.load()
|
||||
|
||||
cfg = load_config(path)
|
||||
names_before = len(store.get_all())
|
||||
prof = math_latex_profile(name="MaesProfil")
|
||||
store.upsert(prof)
|
||||
store.set_active("MaesProfil")
|
||||
assert store.active().name == "MaesProfil"
|
||||
assert len(store.get_all()) == names_before + 1
|
||||
|
||||
assert cfg.provider == "ollama"
|
||||
assert cfg.hotkey == "ctrl+alt+a"
|
||||
assert cfg.math_text_format == "plain"
|
||||
assert "Réponds directement" in cfg.system_prompt
|
||||
prof2 = Profile(name="MaesProfil", provider="openai", model="gpt-4o-mini")
|
||||
store.upsert(prof2)
|
||||
assert store.get("MaesProfil").provider == "openai"
|
||||
|
||||
store.remove("MaesProfil")
|
||||
with pytest.raises(KeyError):
|
||||
store.get("MaesProfil")
|
||||
|
||||
|
||||
def test_accept_unicode_format(tmp_path):
|
||||
path = tmp_path / "config.json"
|
||||
path.write_text(json.dumps({"provider": "ollama", "math_text_format": "unicode"}), encoding="utf-8")
|
||||
|
||||
cfg = load_config(path)
|
||||
|
||||
assert cfg.math_text_format == "unicode"
|
||||
def test_remove_last_profile_blocked(tmp_path):
|
||||
store = ConfigStore(tmp_path / "config.json")
|
||||
store.ensure_defaults()
|
||||
store.load()
|
||||
# Il y a 2 profils par défaut : le premier retrait réussit…
|
||||
store.remove(store.get_all()[0].name)
|
||||
assert len(store.get_all()) == 1
|
||||
# …mais retirer le dernier est interdit.
|
||||
with pytest.raises(ConfigError):
|
||||
store.remove(store.get_all()[0].name)
|
||||
|
||||
|
||||
def test_timeout_zero_disables_timeout(tmp_path):
|
||||
path = tmp_path / "config.json"
|
||||
path.write_text(json.dumps({"provider": "ollama", "request_timeout_seconds": 0}), encoding="utf-8")
|
||||
|
||||
cfg = load_config(path)
|
||||
|
||||
assert cfg.request_timeout_seconds is None
|
||||
def test_set_active_unknown_raises(tmp_path):
|
||||
store = ConfigStore(tmp_path / "config.json")
|
||||
store.ensure_defaults()
|
||||
store.load()
|
||||
with pytest.raises(KeyError):
|
||||
store.set_active("inexistant")
|
||||
|
||||
|
||||
def test_reject_invalid_math_text_format(tmp_path):
|
||||
path = tmp_path / "config.json"
|
||||
path.write_text(json.dumps({"provider": "ollama", "math_text_format": "bad"}), encoding="utf-8")
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
load_config(path)
|
||||
|
||||
|
||||
def test_reject_invalid_provider(tmp_path):
|
||||
path = tmp_path / "config.json"
|
||||
path.write_text(json.dumps({"provider": "bad"}), encoding="utf-8")
|
||||
|
||||
with pytest.raises(ValueError):
|
||||
load_config(path)
|
||||
def test_load_config_return_store(tmp_path):
|
||||
store = load_config(str(tmp_path / "config.json"))
|
||||
assert isinstance(store, ConfigStore)
|
||||
@@ -0,0 +1,97 @@
|
||||
"""Tests du stockage sécurisé (keyring mocké pour éviter tout accès au
|
||||
gestionnaire d'identifiants réel de la machine)."""
|
||||
|
||||
import pytest
|
||||
|
||||
from ai_typewriter.credentials import CredentialError, SecureStore, get_cred
|
||||
|
||||
|
||||
class FakeKeyring:
|
||||
"""Mini stub de keyring en mémoire, exposant la même interface."""
|
||||
|
||||
_data = {}
|
||||
|
||||
@classmethod
|
||||
def reset(cls):
|
||||
cls._data = {}
|
||||
|
||||
@classmethod
|
||||
def set_password(cls, service, username, password):
|
||||
cls._data[(service, username)] = password
|
||||
|
||||
@classmethod
|
||||
def get_password(cls, service, username):
|
||||
return cls._data.get((service, username))
|
||||
|
||||
@classmethod
|
||||
def delete_password(cls, service, username):
|
||||
cls._data.pop((service, username), None)
|
||||
|
||||
|
||||
class FakeKeyringErrors:
|
||||
class PasswordDeleteError(RuntimeError):
|
||||
pass
|
||||
|
||||
|
||||
def test_store_and_get(monkeypatch):
|
||||
import ai_typewriter.credentials as cred
|
||||
|
||||
monkeypatch.setattr(cred, "keyring", FakeKeyring)
|
||||
monkeypatch.setattr(cred.keyring, "errors", FakeKeyringErrors, raising=False)
|
||||
FakeKeyring.reset()
|
||||
|
||||
store = SecureStore("test-service")
|
||||
store.store("openai", "sk-secret")
|
||||
assert store.get("openai") == "sk-secret"
|
||||
assert store.has("openai") is True
|
||||
|
||||
|
||||
def test_empty_credential_returns_empty(monkeypatch):
|
||||
import ai_typewriter.credentials as cred
|
||||
|
||||
monkeypatch.setattr(cred, "keyring", FakeKeyring)
|
||||
FakeKeyring.reset()
|
||||
|
||||
store = SecureStore()
|
||||
assert store.get("") == ""
|
||||
assert store.has("") is False
|
||||
|
||||
|
||||
def test_get_cred_helper_creates_store(monkeypatch):
|
||||
import ai_typewriter.credentials as cred
|
||||
|
||||
monkeypatch.setattr(cred, "keyring", FakeKeyring)
|
||||
monkeypatch.setattr(cred.keyring, "errors", FakeKeyringErrors, raising=False)
|
||||
FakeKeyring.reset()
|
||||
|
||||
# via un store passé explicitement
|
||||
store = SecureStore("t")
|
||||
store.store("k", "v")
|
||||
assert get_cred(store, "k") == "v"
|
||||
# via None (crée un store par défaut, mais retombe sur service réel) —
|
||||
# on vérifie que quelques accesseurs ne plantent pas.
|
||||
assert get_cred(None, "") == ""
|
||||
|
||||
|
||||
def test_store_invalid_credential_name(monkeypatch):
|
||||
import ai_typewriter.credentials as cred
|
||||
|
||||
monkeypatch.setattr(cred, "keyring", FakeKeyring)
|
||||
FakeKeyring.reset()
|
||||
|
||||
store = SecureStore()
|
||||
with pytest.raises(CredentialError):
|
||||
store.store("", "secret")
|
||||
|
||||
|
||||
def test_delete(monkeypatch):
|
||||
import ai_typewriter.credentials as cred
|
||||
|
||||
monkeypatch.setattr(cred, "keyring", FakeKeyring)
|
||||
monkeypatch.setattr(cred.keyring, "errors", FakeKeyringErrors, raising=False)
|
||||
FakeKeyring.reset()
|
||||
|
||||
store = SecureStore("t")
|
||||
store.store("gemini", "cle")
|
||||
store.delete("gemini")
|
||||
assert store.has("gemini") is False
|
||||
@@ -0,0 +1,106 @@
|
||||
"""Tests du moteur : capture -> IA -> dactylographie (tout mocké)."""
|
||||
|
||||
from types import SimpleNamespace
|
||||
|
||||
from ai_typewriter.config import ConfigStore
|
||||
from ai_typewriter.engine import AITypewriterEngine
|
||||
|
||||
|
||||
def _make_store(tmp_path) -> ConfigStore:
|
||||
store = ConfigStore(tmp_path / "config.json")
|
||||
store.ensure_defaults()
|
||||
store.load()
|
||||
return store
|
||||
|
||||
|
||||
def _fake_keyboard(monkeypatch, typed, sent):
|
||||
hook = {}
|
||||
monkeypatch.setattr(
|
||||
"ai_typewriter.key_stepper.keyboard.write",
|
||||
lambda char, delay=0, exact=True: typed.append(char),
|
||||
)
|
||||
monkeypatch.setattr("ai_typewriter.key_stepper.keyboard.send", lambda v: sent.append(v))
|
||||
monkeypatch.setattr(
|
||||
"ai_typewriter.key_stepper.keyboard.hook",
|
||||
lambda callback, suppress=True: hook.update({"callback": callback}) or "hook",
|
||||
)
|
||||
monkeypatch.setattr("ai_typewriter.key_stepper.keyboard.unhook", lambda v: None)
|
||||
return hook
|
||||
|
||||
|
||||
def test_capture_ask_and_step_types_response(monkeypatch, tmp_path):
|
||||
store = _make_store(tmp_path)
|
||||
engine = AITypewriterEngine(store)
|
||||
typed, sent = [], []
|
||||
hook = _fake_keyboard(monkeypatch, typed, sent)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"ai_typewriter.engine.capture_clipboard", lambda: "Question de test"
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"ai_typewriter.engine.ask_ai", lambda prompt, profile, store: "Réponse IA"
|
||||
)
|
||||
|
||||
engine.capture_ask_and_step()
|
||||
assert engine._stepper is not None
|
||||
hook["callback"](SimpleNamespace(event_type="down"))
|
||||
hook["callback"](SimpleNamespace(event_type="down"))
|
||||
assert typed == ["R", "é"]
|
||||
|
||||
|
||||
def test_capture_ask_and_step_with_equations(monkeypatch, tmp_path):
|
||||
store = _make_store(tmp_path)
|
||||
store.set_active("Mathématiques (LaTeX)")
|
||||
engine = AITypewriterEngine(store)
|
||||
typed, sent = [], []
|
||||
hook = _fake_keyboard(monkeypatch, typed, sent)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"ai_typewriter.engine.capture_clipboard", lambda: "Calcule"
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"ai_typewriter.engine.ask_ai", lambda prompt, profile, store: r"[EQ]a=x[/EQ]"
|
||||
)
|
||||
|
||||
engine.capture_ask_and_step()
|
||||
# 1:[EQ]->alt+= ; a: char ; =: char ; x: char ; [/EQ]->right => 5 actions
|
||||
for _ in range(5):
|
||||
hook["callback"](SimpleNamespace(event_type="down"))
|
||||
|
||||
assert sent == ["alt+=", "right"]
|
||||
assert typed == ["a", "=", "x"]
|
||||
|
||||
|
||||
def test_empty_clipboard_does_not_ask(monkeypatch, tmp_path):
|
||||
store = _make_store(tmp_path)
|
||||
engine = AITypewriterEngine(store)
|
||||
called = []
|
||||
monkeypatch.setattr("ai_typewriter.engine.capture_clipboard", lambda: " ")
|
||||
monkeypatch.setattr(
|
||||
"ai_typewriter.engine.ask_ai",
|
||||
lambda prompt, profile, store: called.append(prompt) or "x",
|
||||
)
|
||||
engine.capture_ask_and_step()
|
||||
assert called == []
|
||||
|
||||
|
||||
def test_ask_only_returns_answer(monkeypatch, tmp_path):
|
||||
store = _make_store(tmp_path)
|
||||
engine = AITypewriterEngine(store)
|
||||
monkeypatch.setattr(
|
||||
"ai_typewriter.engine.ask_ai",
|
||||
lambda prompt, profile, store: "reponse-brute",
|
||||
)
|
||||
assert engine.ask_only("bonjour") == "reponse-brute"
|
||||
|
||||
|
||||
def test_concurrent_hotkey_ignored_while_busy(monkeypatch, tmp_path):
|
||||
store = _make_store(tmp_path)
|
||||
engine = AITypewriterEngine(store)
|
||||
|
||||
# Occupe le verrou
|
||||
assert engine._busy.acquire(blocking=False) is True
|
||||
calls = []
|
||||
engine.handle_hotkey() # doit être ignoré (busy)
|
||||
assert calls == []
|
||||
engine._busy.release()
|
||||
+130
-42
@@ -1,48 +1,108 @@
|
||||
from types import SimpleNamespace
|
||||
|
||||
from ai_typewriter.key_stepper import KeyStepper
|
||||
from ai_typewriter.key_stepper import Action, KeyStepper, build_actions
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# build_actions : découpage en actions (sans clavier)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def test_plain_text_actions_single_char_per_action():
|
||||
actions = build_actions("abc")
|
||||
assert actions == [
|
||||
Action("char", "a"),
|
||||
Action("char", "b"),
|
||||
Action("char", "c"),
|
||||
]
|
||||
|
||||
|
||||
def test_equation_markers_replaced_by_key_sequences():
|
||||
actions = build_actions(
|
||||
r"Les racines sont [EQ]z_1 = x[/EQ].",
|
||||
equation_enabled=True,
|
||||
)
|
||||
assert Action("seq", "alt+=") in actions
|
||||
assert Action("seq", "right") in actions
|
||||
# aucun marqueur littéral ne doit être tapé
|
||||
assert Action("char", "[") not in actions
|
||||
kinds = [a.kind for a in actions]
|
||||
assert kinds.count("seq") == 2
|
||||
# le contenu LaTeX est conservé caractère par caractère
|
||||
chars = "".join(a.value for a in actions if a.kind == "char")
|
||||
assert "z_1 = x" in chars
|
||||
|
||||
|
||||
def test_multiple_equations_each_get_start_and_end():
|
||||
text = r"[EQ]a[/EQ] et [EQ]b[/EQ]"
|
||||
actions = build_actions(text, equation_enabled=True)
|
||||
seqs = [a.value for a in actions if a.kind == "seq"]
|
||||
assert seqs == ["alt+=", "right", "alt+=", "right"]
|
||||
|
||||
|
||||
def test_equation_disabled_keeps_markers_literal():
|
||||
text = r"[EQ]\frac{a}{b}[/EQ]"
|
||||
assert build_actions(text, equation_enabled=False) == [
|
||||
Action("char", c) for c in text
|
||||
]
|
||||
|
||||
|
||||
def test_custom_markers_and_keys():
|
||||
actions = build_actions(
|
||||
"<<START>>x^2<<END>>",
|
||||
equation_enabled=True,
|
||||
eq_start_marker="<<START>>",
|
||||
eq_end_marker="<<END>>",
|
||||
eq_start_key="ctrl+shift+e",
|
||||
eq_end_key="space",
|
||||
)
|
||||
assert Action("seq", "ctrl+shift+e") in actions
|
||||
assert Action("seq", "space") in actions
|
||||
# marqueurs retirés, non tapés
|
||||
chars = "".join(a.value for a in actions if a.kind == "char")
|
||||
assert chars == "x^2"
|
||||
assert "START" not in chars
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# KeyStepper (monkeypatch du hook clavier)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _install_hook(monkeypatch, typed, sent, unhooked):
|
||||
hook = {}
|
||||
monkeypatch.setattr(
|
||||
"ai_typewriter.key_stepper.keyboard.write",
|
||||
lambda char, delay=0, exact=True: typed.append(char),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"ai_typewriter.key_stepper.keyboard.send",
|
||||
lambda value: sent.append(value),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"ai_typewriter.key_stepper.keyboard.hook",
|
||||
lambda callback, suppress=True: hook.update({"callback": callback}) or "hook",
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
"ai_typewriter.key_stepper.keyboard.unhook", lambda value: unhooked.append(value)
|
||||
)
|
||||
return hook
|
||||
|
||||
|
||||
def test_start_installs_hook_without_typing_immediately(monkeypatch):
|
||||
typed = []
|
||||
hooked = []
|
||||
|
||||
monkeypatch.setattr("ai_typewriter.key_stepper.keyboard.write", lambda char, delay=0, exact=True: typed.append(char))
|
||||
monkeypatch.setattr("ai_typewriter.key_stepper.keyboard.hook", lambda callback, suppress=True: hooked.append((callback, suppress)) or "hook")
|
||||
typed, sent, unhooked = [], [], []
|
||||
_install_hook(monkeypatch, typed, sent, unhooked)
|
||||
|
||||
stepper = KeyStepper("abc")
|
||||
stepper.start()
|
||||
|
||||
assert typed == []
|
||||
assert hooked and hooked[0][1] is True
|
||||
assert typed == [] and sent == []
|
||||
assert stepper.remaining_characters == 3
|
||||
|
||||
|
||||
def test_single_character_response_waits_for_key_then_unhooks(monkeypatch):
|
||||
typed = []
|
||||
hook = {}
|
||||
unhooked = []
|
||||
|
||||
monkeypatch.setattr("ai_typewriter.key_stepper.keyboard.write", lambda char, delay=0, exact=True: typed.append(char))
|
||||
monkeypatch.setattr("ai_typewriter.key_stepper.keyboard.hook", lambda callback, suppress=True: hook.update({"callback": callback}) or "hook")
|
||||
monkeypatch.setattr("ai_typewriter.key_stepper.keyboard.unhook", lambda value: unhooked.append(value))
|
||||
|
||||
stepper = KeyStepper("x")
|
||||
stepper.start()
|
||||
hook["callback"](SimpleNamespace(event_type="down"))
|
||||
|
||||
assert typed == ["x"]
|
||||
assert unhooked == ["hook"]
|
||||
assert stepper.remaining_characters == 0
|
||||
|
||||
|
||||
def test_each_key_down_types_next_character(monkeypatch):
|
||||
typed = []
|
||||
hook = {}
|
||||
|
||||
monkeypatch.setattr("ai_typewriter.key_stepper.keyboard.write", lambda char, delay=0, exact=True: typed.append(char))
|
||||
monkeypatch.setattr("ai_typewriter.key_stepper.keyboard.hook", lambda callback, suppress=True: hook.update({"callback": callback}) or "hook")
|
||||
monkeypatch.setattr("ai_typewriter.key_stepper.keyboard.unhook", lambda value: None)
|
||||
typed, sent, unhooked = [], [], []
|
||||
hook = _install_hook(monkeypatch, typed, sent, unhooked)
|
||||
|
||||
stepper = KeyStepper("ab")
|
||||
stepper.start()
|
||||
@@ -51,19 +111,47 @@ def test_each_key_down_types_next_character(monkeypatch):
|
||||
|
||||
assert typed == ["a", "b"]
|
||||
assert stepper.remaining_characters == 0
|
||||
assert unhooked # libère le hook à la fin
|
||||
|
||||
|
||||
def test_latex_markers_are_typed_literally_character_by_character(monkeypatch):
|
||||
typed = []
|
||||
hook = {}
|
||||
def test_single_char_response_waits_for_key_then_unhooks(monkeypatch):
|
||||
typed, sent, unhooked = [], [], []
|
||||
hook = _install_hook(monkeypatch, typed, sent, unhooked)
|
||||
|
||||
monkeypatch.setattr("ai_typewriter.key_stepper.keyboard.write", lambda char, delay=0, exact=True: typed.append(char))
|
||||
monkeypatch.setattr("ai_typewriter.key_stepper.keyboard.hook", lambda callback, suppress=True: hook.update({"callback": callback}) or "hook")
|
||||
monkeypatch.setattr("ai_typewriter.key_stepper.keyboard.unhook", lambda value: None)
|
||||
|
||||
stepper = KeyStepper(r"[EQ]\frac{a}{b}[/EQ]")
|
||||
stepper = KeyStepper("x")
|
||||
stepper.start()
|
||||
for _ in range(len(stepper.text)):
|
||||
hook["callback"](SimpleNamespace(event_type="down"))
|
||||
|
||||
assert typed == ["x"] and sent == []
|
||||
assert unhooked == ["hook"]
|
||||
|
||||
|
||||
def test_equation_markers_trigger_key_sequences_while_typing_latex(monkeypatch):
|
||||
typed, sent, unhooked = [], [], []
|
||||
hook = _install_hook(monkeypatch, typed, sent, unhooked)
|
||||
|
||||
stepper = KeyStepper(
|
||||
r"[EQ]a=x[/EQ]",
|
||||
equation_enabled=True,
|
||||
eq_start_key="alt+=",
|
||||
eq_end_key="right",
|
||||
)
|
||||
stepper.start()
|
||||
# 1: [EQ] -> Alt+= ; 2-4: a, =, x ; 5: [/EQ] -> right
|
||||
for _ in range(5):
|
||||
hook["callback"](SimpleNamespace(event_type="down"))
|
||||
|
||||
assert "".join(typed) == r"[EQ]\frac{a}{b}[/EQ]"
|
||||
assert sent == ["alt+=", "right"]
|
||||
assert typed == ["a", "=", "x"]
|
||||
assert stepper.remaining_characters == 0
|
||||
|
||||
|
||||
def test_empty_text_does_not_install_hook(monkeypatch):
|
||||
called = []
|
||||
monkeypatch.setattr(
|
||||
"ai_typewriter.key_stepper.keyboard.hook",
|
||||
lambda callback, suppress=True: called.append(1),
|
||||
)
|
||||
stepper = KeyStepper("")
|
||||
stepper.start()
|
||||
assert called == []
|
||||
@@ -1,18 +0,0 @@
|
||||
import pytest
|
||||
|
||||
from ai_typewriter.math_format import format_math_text
|
||||
|
||||
|
||||
def test_plain_mode_keeps_latex_and_markers_literal():
|
||||
assert format_math_text(r"[EQ]\frac{a}{b}[/EQ]", "plain") == r"[EQ]\frac{a}{b}[/EQ]"
|
||||
|
||||
|
||||
def test_unicode_mode_converts_indices_exponents_and_symbols():
|
||||
result = format_math_text(r"z_1 + x^2 + x_(i+1) + \alpha + \infty", "unicode")
|
||||
|
||||
assert result == "z₁ + x² + xᵢ₊₁ + α + ∞"
|
||||
|
||||
|
||||
def test_invalid_math_text_format_is_rejected():
|
||||
with pytest.raises(ValueError):
|
||||
format_math_text("x_1", "bad")
|
||||
@@ -0,0 +1,90 @@
|
||||
"""Tests de l'actuaire des modèles : listage local, catalogue en ligne et
|
||||
téléchargement (le tout mocké, sans accéder au réseau ni à la CLI ollama)."""
|
||||
|
||||
from ai_typewriter import model_catalog
|
||||
|
||||
|
||||
class FakeResponse:
|
||||
def __init__(self, payload=None, status=200):
|
||||
self.payload = payload
|
||||
self.status_code = status
|
||||
|
||||
def raise_for_status(self):
|
||||
if self.status_code >= 400:
|
||||
raise RuntimeError(f"HTTP {self.status_code}")
|
||||
return None
|
||||
|
||||
def json(self):
|
||||
return self.payload
|
||||
|
||||
|
||||
def test_list_local_models_parses_names(monkeypatch):
|
||||
def fake_get(url, timeout, **kwargs):
|
||||
return FakeResponse(
|
||||
{"models": [{"name": "llama3.1"}, {"model": "mistral"}, {"name": "llama3.1"}]}
|
||||
)
|
||||
|
||||
monkeypatch.setattr("ai_typewriter.model_catalog.requests.get", fake_get)
|
||||
|
||||
names = model_catalog.list_local_models("http://local:11434")
|
||||
assert names == ["llama3.1", "mistral"] # triés + dédupliqués
|
||||
|
||||
|
||||
def test_list_local_models_returns_empty_on_error(monkeypatch):
|
||||
import requests
|
||||
|
||||
def fake_get(url, timeout, **kwargs):
|
||||
raise requests.ConnectionError("boom")
|
||||
|
||||
monkeypatch.setattr("ai_typewriter.model_catalog.requests.get", fake_get)
|
||||
assert model_catalog.list_local_models() == []
|
||||
|
||||
|
||||
def test_search_online_filters_by_query():
|
||||
results = model_catalog.search_online_models("llama")
|
||||
assert results
|
||||
assert all("llama" in m.name for m in results)
|
||||
assert all(m.source == "registry" for m in results)
|
||||
|
||||
|
||||
def test_provider_info():
|
||||
info = model_catalog.provider_info("openai")
|
||||
assert info["needs_key"] is True
|
||||
info_ollama = model_catalog.provider_info("ollama")
|
||||
assert info_ollama["needs_key"] is False
|
||||
|
||||
|
||||
def test_resolve_exact_model_true_when_200(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"ai_typewriter.model_catalog.requests.get", lambda url, timeout: FakeResponse(status=200)
|
||||
)
|
||||
assert model_catalog.resolve_exact_model("llama3.1") is True
|
||||
|
||||
|
||||
def test_resolve_exact_model_false_when_notfound(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"ai_typewriter.model_catalog.requests.get", lambda url, timeout: FakeResponse(status=404)
|
||||
)
|
||||
assert model_catalog.resolve_exact_model("n-existe-pas") is False
|
||||
|
||||
|
||||
def test_pull_model_uses_cli_when_available(monkeypatch):
|
||||
calls = []
|
||||
|
||||
class FakePopen:
|
||||
def __init__(self, cmd, **kwargs):
|
||||
calls.append(cmd)
|
||||
|
||||
monkeypatch.setattr("ai_typewriter.model_catalog.shutil.which", lambda name: "/usr/bin/ollama")
|
||||
monkeypatch.setattr("ai_typewriter.model_catalog.subprocess.Popen", FakePopen)
|
||||
|
||||
model_catalog.pull_model("llama3.1", server_url="http://local")
|
||||
assert calls and calls[0] == ["/usr/bin/ollama", "pull", "llama3.1"]
|
||||
|
||||
|
||||
def test_has_internet_true(monkeypatch):
|
||||
monkeypatch.setattr(
|
||||
"ai_typewriter.model_catalog.requests.head",
|
||||
lambda url, timeout, allow_redirects: FakeResponse(status=200),
|
||||
)
|
||||
assert model_catalog.has_internet() is True
|
||||
Reference in New Issue
Block a user