FloodDiffusion2-Live / cloud_sessions.py
caiyiyi1998's picture
Initial commit
9a25493
Raw History Blame Contribute Delete
3.8 kB
"""File-backed controls shared by the web process and ZeroGPU workers."""
from pathlib import Path
import json
import math
import re
import time
import uuid
from filelock import FileLock
class SessionStore:
def __init__(self, root):
self.root = Path(root)
self.root.mkdir(parents=True, exist_ok=True)
def path(self, session):
if not isinstance(session, str) or not re.fullmatch('[a-f0-9]{32}', session):
raise ValueError('Start a session first.')
return self.root / f'{session}.json'
def _read(self, session, owner):
data = json.loads(self.path(session).read_text())
if not owner or data['owner'] != owner:
raise ValueError('This session belongs to another browser.')
if data.get('stream_active') and time.time() > data.get('lease_expires_at', float('inf')):
data.update(state='paused', stop=True, stream_active=False,
x=0., z=0., keys_updated_at=0.)
self._write(session, data)
return data
def _write(self, session, data):
path = self.path(session)
temporary = path.with_suffix(f'.{uuid.uuid4().hex}.tmp')
temporary.write_text(json.dumps(data))
temporary.replace(path)
def create(self, session, data):
with FileLock(str(self.path(session)) + '.lock'):
if self.path(session).exists():
raise ValueError('This session already exists.')
self._write(session, data)
def read(self, session, owner):
with FileLock(str(self.path(session)) + '.lock'):
return self._read(session, owner)
def patch(self, session, owner, **values):
with FileLock(str(self.path(session)) + '.lock'):
data = self._read(session, owner)
if values.get('state') == 'running' and data['stop']:
values['state'] = 'pausing' if data.get('stream_active') else 'paused'
data.update(values)
self._write(session, data)
return data
def claim(self, session, owner):
with FileLock(str(self.path(session)) + '.lock'):
data = self._read(session, owner)
if data['state'] != 'starting' or data['stop'] or data.get('stream_active'):
raise ValueError('This session is already streaming or has stopped.')
data.update(stream_active=True, lease_expires_at=time.time()+65.)
self._write(session, data)
return data
def control(self, session, owner, seq, x, z, slot):
x, z, selected = float(x), float(z), float(slot)
if not all(math.isfinite(value) for value in (x, z, selected)):
raise ValueError('Invalid direction or prompt slot.')
if selected not in (0., 1., 2., 3.):
raise ValueError('Choose prompt slot 1–4.')
with FileLock(str(self.path(session)) + '.lock'):
data = self._read(session, owner)
if data['stop'] or data['state'] not in ('starting', 'running') or int(seq) <= data['client_seq']:
return data
values = dict(x=max(-1., min(1., x)), z=max(-1., min(1., z)), slot=int(selected))
if any(data[key] != value for key, value in values.items()):
data['version'] += 1
data.update(values, client_seq=int(seq), keys_updated_at=time.time())
self._write(session, data)
return data
def stop(self, session, owner):
with FileLock(str(self.path(session)) + '.lock'):
data = self._read(session, owner)
data.update(stop=True, x=0., z=0., keys_updated_at=0.)
data['state'] = 'pausing' if data.get('stream_active') else 'paused'
self._write(session, data)
return data