Spaces:
Paused
Paused
| import json | |
| import random | |
| import socket | |
| import ssl | |
| import time | |
| from threading import Thread | |
| from curl_cffi import requests | |
| from websocket import WebSocketApp | |
| from .config import DEFAULT_HEADERS, ENDPOINT_SOCKET_IO, LABS_MODELS | |
| from .exceptions import InvalidModelError, NetworkError | |
| from .logger import get_logger | |
| logger = get_logger("labs") | |
| class LabsClient: | |
| def __init__(self): | |
| self.session = requests.Session(headers=DEFAULT_HEADERS.copy(), impersonate="chrome") | |
| self.timestamp = format(random.getrandbits(32), "08x") | |
| poll_url = f"{ENDPOINT_SOCKET_IO}?EIO=4&transport=polling&t={self.timestamp}" | |
| self.sid = json.loads(self.session.get(poll_url).text[1:])["sid"] | |
| self.last_answer = None | |
| self.history = [] | |
| auth_url = ( | |
| f"{ENDPOINT_SOCKET_IO}?EIO=4&transport=polling" | |
| f"&t={self.timestamp}&sid={self.sid}" | |
| ) | |
| assert self.session.post(auth_url, data='40{"jwt":"anonymous-ask-user"}').text == "OK" | |
| context = ssl.create_default_context() | |
| context.minimum_version = ssl.TLSVersion.TLSv1_3 | |
| self.sock = context.wrap_socket( | |
| socket.create_connection(("www.perplexity.ai", 443)), | |
| server_hostname="www.perplexity.ai", | |
| ) | |
| websocket_url = ( | |
| "wss://www.perplexity.ai/socket.io/?EIO=4&transport=websocket" | |
| f"&sid={self.sid}" | |
| ) | |
| self.ws = WebSocketApp( | |
| url=websocket_url, | |
| header={"User-Agent": self.session.headers["User-Agent"]}, | |
| cookie="; ".join( | |
| [f"{key}={value}" for key, value in self.session.cookies.get_dict().items()] | |
| ), | |
| on_open=lambda ws: (ws.send("2probe"), ws.send("5")), | |
| on_message=self._on_message, | |
| on_error=lambda ws, error: logger.error(f"WebSocket Error: {error}"), | |
| socket=self.sock, | |
| ) | |
| Thread(target=self.ws.run_forever, daemon=True).start() | |
| while not (self.ws.sock and self.ws.sock.connected): | |
| time.sleep(0.01) | |
| def _on_message(self, ws, message): | |
| try: | |
| if message == "2": | |
| ws.send("3") | |
| if message.startswith("42"): | |
| response = json.loads(message[2:])[1] | |
| if "final" in response: | |
| self.last_answer = response | |
| except json.JSONDecodeError as e: | |
| logger.error(f"JSON decode error in labs message: {e}") | |
| except Exception as e: | |
| logger.error(f"Error in labs message handler: {e}") | |
| def ask(self, query, model="r1-1776", stream=False): | |
| if model not in LABS_MODELS: | |
| raise InvalidModelError( | |
| f"Invalid labs model '{model}'. Must be one of: {', '.join(LABS_MODELS)}" | |
| ) | |
| self.last_answer = None | |
| self.history.append({"role": "user", "content": query}) | |
| self.ws.send( | |
| "42" + json.dumps( | |
| [ | |
| "perplexity_labs", | |
| { | |
| "messages": self.history, | |
| "model": model, | |
| "source": "default", | |
| "version": "2.18", | |
| }, | |
| ] | |
| ) | |
| ) | |
| def stream_response(): | |
| answer = None | |
| while True: | |
| if self.last_answer != answer: | |
| answer = self.last_answer | |
| yield answer | |
| if self.last_answer and self.last_answer.get("final"): | |
| answer = self.last_answer | |
| self.last_answer = None | |
| self.history.append( | |
| { | |
| "role": "assistant", | |
| "content": answer["output"], | |
| "priority": 0, | |
| } | |
| ) | |
| return | |
| time.sleep(0.01) | |
| if stream: | |
| return stream_response() | |
| while True: | |
| if self.last_answer and self.last_answer.get("final"): | |
| answer = self.last_answer | |
| self.last_answer = None | |
| self.history.append( | |
| { | |
| "role": "assistant", | |
| "content": answer["output"], | |
| "priority": 0, | |
| } | |
| ) | |
| return answer | |
| time.sleep(0.01) | |