| 123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191 |
- 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]
|