Aller au contenu principal

Implémenter le streaming en Python

Du prototype à la production

La leçon précédente a introduit les concepts du streaming SSE. Passons maintenant à l’implémentation concrète : un client de streaming robuste, avec gestion d’erreurs, accumulation du texte et intégration dans une application réelle.

Le pattern de base

Voici le squelette minimal que vous adapterez à chaque projet :

from mistralai import Mistral
import os

client = Mistral(api_key=os.getenv("MISTRAL_API_KEY"))

def stream_chat(messages: list, model: str = "mistral-large-latest") -> str:
    """Envoie une requête en streaming et affiche les tokens en temps réel."""
    stream = client.chat.stream(
        model=model,
        messages=messages
    )

    full_text = ""
    for chunk in stream:
        content = chunk.data.choices[0].delta.content
        if content:
            full_text += content
            print(content, end="", flush=True)

    print()  # Retour à la ligne
    return full_text

Gestion complète des erreurs

En production, les erreurs réseau, les timeouts et les limites de débit sont inévitables. Voici un client robuste :

from mistralai import Mistral
from mistralai.exceptions import MistralAPIException
import os
import time

client = Mistral(api_key=os.getenv("MISTRAL_API_KEY"))

def stream_with_retry(
    messages: list,
    model: str = "mistral-large-latest",
    max_retries: int = 3,
    max_tokens: int = 2000
) -> dict:
    """Streaming avec retry automatique et métriques."""
    for attempt in range(max_retries):
        try:
            stream = client.chat.stream(
                model=model,
                messages=messages,
                max_tokens=max_tokens
            )

            full_text = ""
            finish_reason = None
            chunk_count = 0

            for chunk in stream:
                delta = chunk.data.choices[0].delta
                if delta.content:
                    full_text += delta.content
                    chunk_count += 1
                    yield {"type": "token", "content": delta.content}

                if chunk.data.choices[0].finish_reason:
                    finish_reason = chunk.data.choices[0].finish_reason

            yield {
                "type": "done",
                "full_text": full_text,
                "finish_reason": finish_reason,
                "chunks": chunk_count
            }
            return

        except MistralAPIException as e:
            if e.status_code == 429 and attempt < max_retries - 1:
                wait_time = 2 ** attempt  # Backoff exponentiel
                yield {"type": "retry", "wait": wait_time, "attempt": attempt + 1}
                time.sleep(wait_time)
            else:
                yield {"type": "error", "message": str(e), "status": e.status_code}
                return

        except Exception as e:
            yield {"type": "error", "message": str(e), "status": None}
            return

Utilisation du client avec retry

messages = [
    {"role": "system", "content": "Vous êtes un assistant technique."},
    {"role": "user", "content": "Expliquez les design patterns en Python."}
]

for event in stream_with_retry(messages):
    if event["type"] == "token":
        print(event["content"], end="", flush=True)
    elif event["type"] == "retry":
        print(f"\n[Retry {event['attempt']}, attente {event['wait']}s...]")
    elif event["type"] == "done":
        print(f"\n\n--- Terminé ({event['chunks']} chunks, {event['finish_reason']})")
    elif event["type"] == "error":
        print(f"\n[Erreur : {event['message']}]")

Streaming asynchrone avec asyncio

Pour les applications web ou les serveurs à haute concurrence, utilisez le client asynchrone :

from mistralai import Mistral
import asyncio
import os

client = Mistral(api_key=os.getenv("MISTRAL_API_KEY"))

async def async_stream_chat(messages: list) -> str:
    """Streaming asynchrone pour les applications concurrentes."""
    stream = await client.chat.stream_async(
        model="mistral-large-latest",
        messages=messages
    )

    full_text = ""
    async for chunk in stream:
        content = chunk.data.choices[0].delta.content
        if content:
            full_text += content
            print(content, end="", flush=True)

    print()
    return full_text

# Exécution
async def main():
    messages = [{"role": "user", "content": "Bonjour, présentez-vous."}]
    result = await async_stream_chat(messages)
    print(f"\nTotal : {len(result)} caractères")

asyncio.run(main())

Intégration avec FastAPI

Voici comment exposer le streaming via une API web avec FastAPI :

from fastapi import FastAPI
from fastapi.responses import StreamingResponse
from mistralai import Mistral
import os
import json

app = FastAPI()
client = Mistral(api_key=os.getenv("MISTRAL_API_KEY"))

async def generate_stream(prompt: str):
    """Générateur SSE pour FastAPI."""
    stream = await client.chat.stream_async(
        model="mistral-large-latest",
        messages=[{"role": "user", "content": prompt}]
    )

    async for chunk in stream:
        content = chunk.data.choices[0].delta.content
        if content:
            data = json.dumps({"content": content})
            yield f"data: {data}\n\n"

    yield "data: [DONE]\n\n"

@app.get("/chat/stream")
async def chat_stream(prompt: str):
    return StreamingResponse(
        generate_stream(prompt),
        media_type="text/event-stream"
    )

Mesurer la performance du streaming

Deux métriques clés à surveiller :

import time

start = time.perf_counter()
first_token_time = None
token_count = 0

stream = client.chat.stream(
    model="mistral-large-latest",
    messages=[{"role": "user", "content": "Écrivez un paragraphe sur Paris."}]
)

for chunk in stream:
    if chunk.data.choices[0].delta.content:
        if first_token_time is None:
            first_token_time = time.perf_counter() - start
        token_count += 1

total_time = time.perf_counter() - start

print(f"Time to First Token (TTFT) : {first_token_time:.3f}s")
print(f"Tokens générés             : {token_count}")
print(f"Temps total                : {total_time:.3f}s")
print(f"Tokens/seconde             : {token_count / total_time:.1f}")

Le TTFT (Time to First Token) est la métrique la plus importante pour l’expérience utilisateur. C’est le temps entre l’envoi de la requête et l’affichage du premier caractère.

Points clés à retenir

  • Utilisez un pattern générateur (yield) pour propager les tokens en streaming
  • Implémentez un backoff exponentiel pour les erreurs 429 (rate limit)
  • Préférez stream_async pour les applications web concurrentes
  • Mesurez le TTFT et les tokens/seconde pour optimiser la performance
  • Avec FastAPI, StreamingResponse + text/event-stream propagent le SSE au navigateur