{
 "cells": [
  {
   "cell_type": "markdown",
   "id": "0",
   "metadata": {},
   "source": [
    "# WebSocket Target\n",
    "\n",
    "`WebsocketTarget` connects PyRIT to services that use a custom WebSocket protocol.\n",
    "Supply the service-specific initialization messages, prompt builder, and response parser.\n",
    "The `protocol_identifier` is a non-secret name for this complete protocol configuration.\n",
    "\n",
    "This example starts a local PyRIT WebSocket service. It exercises the real WebSocket\n",
    "transport without credentials or an external endpoint."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "1",
   "metadata": {},
   "outputs": [],
   "source": [
    "import json\n",
    "import uuid\n",
    "\n",
    "from websockets.asyncio.client import ClientConnection\n",
    "from websockets.asyncio.server import ServerConnection, serve\n",
    "\n",
    "from pyrit.models import Message, MessagePiece\n",
    "from pyrit.prompt_target import WebsocketTarget\n",
    "from pyrit.setup import IN_MEMORY, initialize_pyrit_async\n",
    "\n",
    "await initialize_pyrit_async(memory_db_type=IN_MEMORY, load_defaults=False, silent=True)  # type: ignore\n",
    "\n",
    "\n",
    "async def pyrit_websocket_handler(websocket: ServerConnection) -> None:\n",
    "    initialization = json.loads(await websocket.recv())\n",
    "    if initialization != {\"type\": \"initialize\", \"client\": \"PyRIT\"}:\n",
    "        await websocket.close(code=1002, reason=\"Invalid initialization message\")\n",
    "        return\n",
    "\n",
    "    await websocket.send(json.dumps({\"message\": \"PyRIT WebSocket target ready\"}))\n",
    "\n",
    "    async for raw_message in websocket:\n",
    "        request = json.loads(raw_message)\n",
    "        if request[\"type\"] == \"restore\":\n",
    "            await websocket.send(json.dumps({\"type\": \"restored\"}))\n",
    "            continue\n",
    "\n",
    "        await websocket.send(json.dumps({\"event\": \"processing\"}))\n",
    "        await websocket.send(json.dumps({\"message\": f\"PyRIT received: {request['prompt']}\"}))\n",
    "\n",
    "\n",
    "def response_parser(message: str | bytes) -> str | None:\n",
    "    if isinstance(message, bytes):\n",
    "        message = message.decode()\n",
    "    return json.loads(message).get(\"message\")\n",
    "\n",
    "\n",
    "def message_builder(prompt: str) -> str:\n",
    "    return json.dumps({\"type\": \"prompt\", \"prompt\": prompt})\n",
    "\n",
    "\n",
    "async def restore_conversation_async(\n",
    "    websocket: ClientConnection,\n",
    "    conversation_history: list[Message],\n",
    ") -> None:\n",
    "    history = [\n",
    "        {\n",
    "            \"role\": message.message_pieces[0].role,\n",
    "            \"content\": message.get_value(),\n",
    "        }\n",
    "        for message in conversation_history\n",
    "    ]\n",
    "    await websocket.send(json.dumps({\"type\": \"restore\", \"history\": history}))\n",
    "    acknowledgement = json.loads(await websocket.recv())\n",
    "    if acknowledgement != {\"type\": \"restored\"}:\n",
    "        raise ConnectionError(\"The WebSocket service did not restore the conversation.\")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "2",
   "metadata": {},
   "source": [
    "Start the local service on an available loopback port, then configure the target for its protocol.\n",
    "\n",
    "The restore callback is service-specific. PyRIT calls it when a multi-turn conversation needs a\n",
    "replacement connection. Without this callback, the target fails instead of silently losing history."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "3",
   "metadata": {},
   "outputs": [],
   "source": [
    "server = await serve(pyrit_websocket_handler, \"127.0.0.1\", 0)  # type: ignore\n",
    "port = server.sockets[0].getsockname()[1]\n",
    "\n",
    "target = WebsocketTarget(\n",
    "    endpoint=f\"ws://127.0.0.1:{port}\",\n",
    "    protocol_identifier=\"local-pyrit-echo-v1\",\n",
    "    initialization_strings=[json.dumps({\"type\": \"initialize\", \"client\": \"PyRIT\"})],\n",
    "    response_parser=response_parser,\n",
    "    message_builder=message_builder,\n",
    "    conversation_restore_callback=restore_conversation_async,\n",
    "    discard_initial_messages=1,\n",
    ")"
   ]
  },
  {
   "cell_type": "markdown",
   "id": "4",
   "metadata": {},
   "source": [
    "Send a prompt through the target and close both sides of the connection.\n",
    "\n",
    "Cleanup is terminal for this target instance. If a connection fails while a prompt is in\n",
    "progress, the target discards that connection and raises the error instead of retrying the prompt."
   ]
  },
  {
   "cell_type": "code",
   "execution_count": null,
   "id": "5",
   "metadata": {},
   "outputs": [
    {
     "name": "stdout",
     "output_type": "stream",
     "text": [
      "PyRIT received: Hello\n"
     ]
    }
   ],
   "source": [
    "request = MessagePiece(\n",
    "    role=\"user\",\n",
    "    original_value=\"Hello\",\n",
    "    original_value_data_type=\"text\",\n",
    "    conversation_id=str(uuid.uuid4()),\n",
    ").to_message()\n",
    "\n",
    "try:\n",
    "    response = await target.send_prompt_async(message=request)  # type: ignore\n",
    "    print(response[0].get_value())\n",
    "finally:\n",
    "    await target.cleanup_target_async()  # type: ignore\n",
    "    server.close()\n",
    "    await server.wait_closed()  # type: ignore"
   ]
  }
 ],
 "metadata": {
  "language_info": {
   "codemirror_mode": {
    "name": "ipython",
    "version": 3
   },
   "file_extension": ".py",
   "mimetype": "text/x-python",
   "name": "python",
   "nbconvert_exporter": "python",
   "pygments_lexer": "ipython3",
   "version": "3.13.5"
  }
 },
 "nbformat": 4,
 "nbformat_minor": 5
}
