from datetime import datetime, timedelta import json import os import socket class InMemoryTokenStore: def __init__(self, refresh_skew_seconds=120): self.refresh_skew_seconds = int(refresh_skew_seconds) self._session = None def save(self, token, expire_time): self._session = { 'token': token, 'expire_time': expire_time, } return self._session def get(self): return self._session def clear(self): self._session = None def is_expiring(self): if not self._session: return True expire_time = self._session.get('expire_time') if not expire_time: return True if isinstance(expire_time, str) and expire_time.endswith('Z'): expire_time = expire_time[:-1] + '+00:00' expires_at = datetime.fromisoformat(str(expire_time)) return expires_at <= datetime.now(expires_at.tzinfo) + timedelta(seconds=self.refresh_skew_seconds) def require_token(self): if not self._session or not self._session.get('token'): raise RuntimeError('mcp token missing') return self._session['token'] class FileTokenStore(InMemoryTokenStore): def __init__(self, path, refresh_skew_seconds=120): super().__init__(refresh_skew_seconds=refresh_skew_seconds) self.path = path self._session = self._read_session() def save(self, token, expire_time): session = super().save(token, expire_time) directory = os.path.dirname(self.path) if directory: os.makedirs(directory, exist_ok=True) with open(self.path, 'w', encoding='utf-8') as handle: json.dump(session, handle, ensure_ascii=False) return session def get(self): if self._session is None and os.path.exists(self.path): self._session = self._read_session() return self._session def clear(self): super().clear() if os.path.exists(self.path): os.remove(self.path) def _read_session(self): if not os.path.exists(self.path): return None with open(self.path, 'r', encoding='utf-8') as handle: return json.load(handle) class RedisTokenStore(InMemoryTokenStore): def __init__(self, client, prefix='fms:mcp:workbuddy:', session_key='', refresh_skew_seconds=120): super().__init__(refresh_skew_seconds=refresh_skew_seconds) session_key = str(session_key or '').strip() if not session_key: raise ValueError('session_key is required for RedisTokenStore') self.client = client self.prefix = str(prefix or 'fms:mcp:workbuddy:') self.session_key = session_key self.key = self.prefix + self.session_key def save(self, token, expire_time): session = super().save(token, expire_time) payload = json.dumps(session, ensure_ascii=False) ttl = self._ttl_seconds(expire_time) self.client.set(self.key, payload, ex=ttl if ttl and ttl > 0 else None) return session def get(self): raw = self.client.get(self.key) if raw is None or raw == '': self._session = None return None if isinstance(raw, bytes): raw = raw.decode('utf-8') self._session = json.loads(raw) return self._session def clear(self): super().clear() self.client.delete(self.key) def is_expiring(self): self.get() return super().is_expiring() @staticmethod def _ttl_seconds(expire_time): if not expire_time: return None value = expire_time[:-1] + '+00:00' if isinstance(expire_time, str) and expire_time.endswith('Z') else expire_time expires_at = datetime.fromisoformat(str(value)) now = datetime.now(expires_at.tzinfo) if expires_at.tzinfo else datetime.now() return max(1, int((expires_at - now).total_seconds())) class RedisSocketClient: def __init__(self, host='127.0.0.1', port=6379, db=0, password='', timeout=5): self.host = host self.port = int(port) self.db = int(db) self.password = password or '' self.timeout = int(timeout) def set(self, key, value, ex=None): command = ['SET', key, value] if ex is not None: command.extend(['EX', int(ex)]) return self._execute(*command) == 'OK' def get(self, key): return self._execute('GET', key) def delete(self, key): return self._execute('DEL', key) def _execute(self, *parts): with socket.create_connection((self.host, self.port), timeout=self.timeout) as sock: stream = sock.makefile('rb') if self.password: self._send(sock, 'AUTH', self.password) self._read_response(stream) if self.db: self._send(sock, 'SELECT', self.db) self._read_response(stream) self._send(sock, *parts) return self._read_response(stream) @staticmethod def _send(sock, *parts): payload = ['*{0}\r\n'.format(len(parts)).encode('utf-8')] for part in parts: if isinstance(part, bytes): data = part else: data = str(part).encode('utf-8') payload.append('${0}\r\n'.format(len(data)).encode('utf-8')) payload.append(data + b'\r\n') sock.sendall(b''.join(payload)) def _read_response(self, stream): prefix = stream.read(1) if not prefix: raise RuntimeError('empty redis response') if prefix == b'+': return self._read_line(stream).decode('utf-8') if prefix == b'-': raise RuntimeError(self._read_line(stream).decode('utf-8')) if prefix == b':': return int(self._read_line(stream)) if prefix == b'$': length = int(self._read_line(stream)) if length == -1: return None data = stream.read(length) stream.read(2) return data.decode('utf-8') if prefix == b'*': length = int(self._read_line(stream)) return [self._read_response(stream) for _ in range(length)] raise RuntimeError('unsupported redis response: {0}'.format(prefix)) @staticmethod def _read_line(stream): line = stream.readline() if not line.endswith(b'\r\n'): raise RuntimeError('invalid redis line') return line[:-2]