166 lines
6.7 KiB
Python
166 lines
6.7 KiB
Python
"""Fixed-upstream ASR gateway: stdlib only, no retry, no body/header logging."""
|
|
from __future__ import annotations
|
|
import hmac
|
|
import http.client
|
|
import json
|
|
import os
|
|
import socket
|
|
import threading
|
|
import time
|
|
from http.server import BaseHTTPRequestHandler, ThreadingHTTPServer
|
|
from pathlib import Path
|
|
|
|
MAX_BODY = 30 * 1024 * 1024 + 65536 # Dify accepts 30 MiB file; multipart allowance.
|
|
MAX_RESPONSE = 4 * 1024 * 1024
|
|
UPSTREAM_HOST = 'followup-audio-asr'
|
|
UPSTREAM_PORT = 8000
|
|
RATE_PER_MINUTE = 120
|
|
RATE_BURST = 4
|
|
|
|
class Gate:
|
|
def __init__(self, rate=RATE_PER_MINUTE, burst=RATE_BURST):
|
|
self.rate, self.burst = rate, burst
|
|
self.tokens, self.at = float(burst), time.monotonic()
|
|
self.lock = threading.Lock()
|
|
self.active = threading.BoundedSemaphore(1)
|
|
|
|
def take(self):
|
|
with self.lock:
|
|
now = time.monotonic()
|
|
self.tokens = min(float(self.burst), self.tokens + (now - self.at) * self.rate / 60)
|
|
self.at = now
|
|
if self.tokens < 1:
|
|
return False
|
|
self.tokens -= 1
|
|
return True
|
|
|
|
class Server(ThreadingHTTPServer):
|
|
daemon_threads = True
|
|
allow_reuse_address = True
|
|
request_queue_size = 8
|
|
|
|
def __init__(self, address, handler, *, key, upstream=(UPSTREAM_HOST, UPSTREAM_PORT), rate=RATE_PER_MINUTE, burst=RATE_BURST, timeout=180):
|
|
super().__init__(address, handler)
|
|
self.key, self.upstream, self.timeout = key, upstream, timeout
|
|
self.gate = Gate(rate, burst)
|
|
self.connections = threading.BoundedSemaphore(8)
|
|
|
|
def process_request(self, request, client_address):
|
|
if not self.connections.acquire(blocking=False):
|
|
try:
|
|
request.sendall(b'HTTP/1.1 429 Too Many Requests\r\nContent-Length: 0\r\nConnection: close\r\n\r\n')
|
|
finally:
|
|
self.shutdown_request(request)
|
|
return
|
|
try:
|
|
super().process_request(request, client_address)
|
|
except BaseException:
|
|
self.connections.release()
|
|
raise
|
|
|
|
def process_request_thread(self, request, client_address):
|
|
try:
|
|
super().process_request_thread(request, client_address)
|
|
finally:
|
|
self.connections.release()
|
|
|
|
def handle_error(self, request, client_address):
|
|
# Deliberately suppress tracebacks/request fragments from client errors.
|
|
print(json.dumps({'event': 'client_error'}), flush=True)
|
|
|
|
class Handler(BaseHTTPRequestHandler):
|
|
protocol_version = 'HTTP/1.1'
|
|
server_version = 'FollowupASRBridge/1'
|
|
sys_version = ''
|
|
|
|
def setup(self):
|
|
super().setup()
|
|
self.connection.settimeout(30)
|
|
|
|
def log_message(self, format, *args):
|
|
pass
|
|
|
|
def reply(self, code, body, content_type='application/json'):
|
|
self.close_connection = True
|
|
self.send_response(code)
|
|
self.send_header('Content-Type', content_type)
|
|
self.send_header('Content-Length', str(len(body)))
|
|
self.send_header('Connection', 'close')
|
|
self.send_header('Cache-Control', 'no-store')
|
|
if code == 429:
|
|
self.send_header('Retry-After', '1')
|
|
self.end_headers()
|
|
self.wfile.write(body)
|
|
|
|
def error_json(self, code, label):
|
|
self.reply(code, json.dumps({'error': label}).encode())
|
|
|
|
def authorized(self):
|
|
supplied = self.headers.get('Authorization', '')
|
|
return hmac.compare_digest(supplied.encode(), b'Bearer ' + self.server.key)
|
|
|
|
def upstream_call(self, method, path, body=None, content_type=None, health=False):
|
|
conn = http.client.HTTPConnection(*self.server.upstream, timeout=3 if health else self.server.timeout)
|
|
try:
|
|
headers = {} if health else {'Authorization': 'Bearer ' + self.server.key.decode()}
|
|
if content_type:
|
|
headers['Content-Type'] = content_type
|
|
conn.request(method, path, body=body, headers=headers)
|
|
response = conn.getresponse()
|
|
result = response.read(MAX_RESPONSE + 1)
|
|
if len(result) > MAX_RESPONSE:
|
|
return self.error_json(502, 'upstream_response_too_large')
|
|
if health:
|
|
return self.reply(200 if response.status == 200 else 503, b'{"status":"ok"}' if response.status == 200 else b'{"status":"unhealthy"}')
|
|
self.reply(response.status, result, response.getheader('Content-Type', 'application/json'))
|
|
except (OSError, http.client.HTTPException, TimeoutError):
|
|
self.error_json(503 if health else 502, 'upstream_unavailable')
|
|
finally:
|
|
conn.close()
|
|
|
|
def do_GET(self):
|
|
if self.path == '/health':
|
|
return self.upstream_call('GET', '/health', health=True)
|
|
if self.path != '/v1/models':
|
|
return self.error_json(404, 'not_found')
|
|
if not self.authorized():
|
|
return self.error_json(401, 'unauthorized')
|
|
return self.upstream_call('GET', '/v1/models')
|
|
|
|
def do_POST(self):
|
|
if self.path != '/v1/audio/transcriptions':
|
|
return self.error_json(404, 'not_found')
|
|
if not self.authorized():
|
|
return self.error_json(401, 'unauthorized')
|
|
if self.headers.get('Transfer-Encoding') or not self.headers.get('Content-Length', '').isdigit():
|
|
return self.error_json(411, 'content_length_required')
|
|
length = int(self.headers['Content-Length'])
|
|
if length > MAX_BODY:
|
|
return self.error_json(413, 'body_too_large')
|
|
if length <= 0:
|
|
return self.error_json(400, 'empty_body')
|
|
content_type = self.headers.get('Content-Type', '')
|
|
if not content_type.startswith('multipart/form-data;') or 'boundary=' not in content_type:
|
|
return self.error_json(415, 'multipart_required')
|
|
if not self.server.gate.active.acquire(blocking=False):
|
|
return self.error_json(429, 'asr_busy')
|
|
try:
|
|
if not self.server.gate.take():
|
|
return self.error_json(429, 'rate_limit')
|
|
try:
|
|
body = self.rfile.read(length)
|
|
except (OSError, socket.timeout):
|
|
return self.error_json(408, 'upload_timeout')
|
|
if len(body) != length:
|
|
return self.error_json(400, 'incomplete_body')
|
|
return self.upstream_call('POST', '/v1/audio/transcriptions', body, content_type)
|
|
finally:
|
|
self.server.gate.active.release()
|
|
|
|
if __name__ == '__main__':
|
|
key = Path('/run/secrets/asr-key').read_bytes().strip()
|
|
if len(key) < 16 or b'\n' in key or b'\r' in key:
|
|
raise SystemExit('invalid_secret_file')
|
|
print(json.dumps({'event': 'started', 'rate_per_minute': RATE_PER_MINUTE, 'inflight': 1, 'host_ports': False}), flush=True)
|
|
Server(('0.0.0.0', 8080), Handler, key=key).serve_forever()
|