同步客户绑定与登录恢复优化

This commit is contained in:
chen qiang
2026-10-08 19:51:38 +08:00
parent 2a81f5668e
commit 3e25ca4696
12 changed files with 896 additions and 152 deletions
+390
View File
@@ -0,0 +1,390 @@
import io
import base64
import json
import os
from pathlib import Path
import subprocess
import sys
import tempfile
import time
import unittest
from contextlib import redirect_stdout, redirect_stderr
from unittest import mock
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from scripts import oracle_skill as skill
from scripts.runtime_state import file_lock, read_config, update_config
def response(status=200, payload=None):
result = mock.Mock(status_code=status, ok=status < 400)
result.json.return_value = payload or {"success": True, "data": "ok"}
if status >= 400:
result.raise_for_status.side_effect = skill.requests.HTTPError(str(status))
return result
def worker(mode, directory, target):
directory = Path(directory)
skill.get_script_dir = lambda: str(directory)
cfg = skill.get_config()
(directory / (target + '.ready')).touch()
deadline = time.monotonic() + 15
while not (directory / 'go').exists():
if time.monotonic() >= deadline:
raise RuntimeError('worker barrier timeout')
time.sleep(.01)
if mode == 'renew':
skill.prepare_command(['skill', 'query', '--client', target, 'SELECT 1 FROM DUAL'])
def renew(values):
with open(directory / 'renewals', 'a') as handle:
handle.write('renew\n')
time.sleep(.2)
values.update(access_token='new-token', expires_at='2999-01-01T00:00:00+00:00')
skill.save_config(values)
return True
skill._auto_device_login_unlocked = renew
rejected = 'old-token' if cfg.get('expires_at', '').startswith('2999') else None
assert skill._auto_device_login(skill.get_config(), rejected)
def post(url, **kwargs):
assert kwargs['headers']['Authorization'] == 'Bearer new-token'
assert kwargs['json']['server_id'] == target
return response()
skill.requests.post = post
assert skill.query(skill.get_server_id(), 'execute_query', sql='SELECT 1 FROM DUAL')['client_code'] == target
elif mode == 'merge':
cfg[target] = target
skill.save_config(cfg)
elif mode == 'crash':
with file_lock(str(directory / 'crash.lock')):
(directory / 'locked').touch()
time.sleep(20)
class RuntimeTests(unittest.TestCase):
def setUp(self):
self.temp = tempfile.TemporaryDirectory()
self.addCleanup(self.temp.cleanup)
self.directory = Path(self.temp.name)
self.path = self.directory / 'config.json'
self.path.write_text(json.dumps({
'transit_url': 'https://transit.example', 'access_token': 'old-token',
'expires_at': '2999-01-01T00:00:00+00:00',
'client_code': 'HENLO', 'server_id': 'HENLO', 'user_name': 'test',
'client_list': [{'code': 'HENLO'}, {'code': 'WEIRUI'}],
}))
patch = mock.patch.object(skill, 'get_script_dir', return_value=str(self.directory))
patch.start()
self.addCleanup(patch.stop)
skill.EXECUTION_CONTEXT = None
skill.JSON_OUTPUT = False
self.addCleanup(self.reset_context)
def reset_context(self):
skill.EXECUTION_CONTEXT = None
skill.JSON_OUTPUT = False
def prepare(self, *args):
return skill.prepare_command(['skill', *args])
def test_explicit_target_does_not_persist_and_survives_default_switch(self):
self.prepare('query', '--client', 'WEIRUI', 'SELECT 1 FROM DUAL')
update_config(str(self.path), {'client_code': 'OTHER', 'server_id': 'OTHER'})
cfg = skill.get_config()
cfg['access_token'] = 'new-token'
skill.save_config(cfg)
self.assertEqual(skill.get_server_id(), 'WEIRUI')
stored = read_config(str(self.path))
self.assertEqual(stored['client_code'], 'OTHER')
self.assertEqual(stored['access_token'], 'new-token')
def test_implicit_default_is_fixed_at_command_start(self):
self.prepare('describe', 'BOSNDS3', 'M_PRODUCT')
update_config(str(self.path), {'client_code': 'WEIRUI', 'server_id': 'WEIRUI', 'transit_url': 'https://other.example'})
self.assertEqual(skill.get_server_id(), 'HENLO')
self.assertEqual(skill.get_transit_url(), 'https://transit.example')
def test_missing_target_stops_before_login(self):
update_config(str(self.path), {'client_code': '', 'server_id': ''})
with mock.patch.object(skill.requests, 'post') as post:
with self.assertRaisesRegex(ValueError, '客户'):
self.prepare('query', 'SELECT 1 FROM DUAL')
post.assert_not_called()
def test_json_missing_target_returns_structured_error(self):
update_config(str(self.path), {'client_code': '', 'server_id': ''})
output = io.StringIO()
with mock.patch.object(sys, 'argv', ['skill', 'query', '--json', 'SELECT 1 FROM DUAL']), redirect_stdout(output), redirect_stderr(io.StringIO()), self.assertRaises(SystemExit):
skill.main()
self.assertFalse(json.loads(output.getvalue())['success'])
def test_empty_explicit_target_and_missing_arguments_are_rejected(self):
for args in [('query', '--client', '', 'SELECT 1 FROM DUAL'), ('describe', '--client', 'HENLO', 'BOSNDS3')]:
with self.subTest(args=args), self.assertRaises(ValueError):
self.prepare(*args)
def test_positional_conflict_and_legacy_layouts(self):
with self.assertRaisesRegex(ValueError, '不一致'):
self.prepare('ops_report', '--client', 'HENLO', 'WEIRUI')
self.assertEqual(self.prepare('ops_report', 'weirui')[-1], 'WEIRUI')
self.assertEqual(self.prepare('awr_download', '--client', 'HENLO', '20261008')[-2:], ['HENLO', '20261008'])
self.assertEqual(self.prepare('agent_update', '--client', 'HENLO', '300')[-2:], ['HENLO', '300'])
def test_sql_is_not_scanned_for_options(self):
sql = "SELECT '--client WEIRUI --json' FROM DUAL"
argv = self.prepare('query', '--client', 'HENLO', sql)
self.assertEqual(argv[-1], sql)
self.assertFalse(skill.JSON_OUTPUT)
def test_permission_lookup_keeps_customer_between_steps(self):
calls = []
def query(server, action, *args, **kwargs):
calls.append(server)
update_config(str(self.path), {'client_code': 'WEIRUI', 'server_id': 'WEIRUI'})
return {'success': True, 'data': '{}'}
with mock.patch.object(skill, 'query', side_effect=query):
skill.query_with_permission('1', '2', 'SELECT 1 FROM DUAL')
self.assertEqual(calls, ['HENLO', 'HENLO'])
def test_valid_token_requires_only_business_request(self):
self.prepare('query', '--client', 'HENLO', 'SELECT 1 FROM DUAL')
with mock.patch.object(skill.requests, 'get') as get, mock.patch.object(skill.requests, 'post', return_value=response()) as post:
self.assertTrue(skill.query('HENLO', 'execute_query', sql='SELECT 1 FROM DUAL')['success'])
get.assert_not_called()
self.assertEqual(post.call_count, 1)
def test_401_recovers_once_and_preserves_payload(self):
self.prepare('query', '--client', 'WEIRUI', 'SELECT 1 FROM DUAL')
def renew(cfg):
cfg['access_token'] = 'new-token'
skill.save_config(cfg)
return True
with mock.patch.object(skill, '_auto_device_login_unlocked', side_effect=renew) as login, mock.patch.object(skill.requests, 'post', side_effect=[response(401), response()]) as post:
result = skill.query('WEIRUI', 'execute_query', sql='SELECT 1 FROM DUAL')
self.assertTrue(result['success'])
self.assertEqual(login.call_count, 1)
self.assertEqual(post.call_count, 2)
self.assertEqual(post.call_args.kwargs['headers']['Authorization'], 'Bearer new-token')
self.assertEqual(post.call_args_list[0].kwargs['json'], post.call_args_list[1].kwargs['json'])
def test_second_401_and_403_do_not_sleep_or_loop(self):
for statuses in ([401, 401], [403]):
with self.subTest(statuses=statuses):
self.reset_context()
self.prepare('query', '--client', 'HENLO', 'SELECT 1 FROM DUAL')
def renew(cfg):
cfg['access_token'] += '-new'
skill.save_config(cfg)
return True
with mock.patch.object(skill, '_auto_device_login_unlocked', side_effect=renew), mock.patch.object(skill.requests, 'post', side_effect=[response(status) for status in statuses]) as post, mock.patch.object(skill.time, 'sleep') as sleep:
self.assertFalse(skill.query('HENLO', 'execute_query')['success'])
self.assertEqual(post.call_count, len(statuses))
sleep.assert_not_called()
def test_rejected_old_token_reuses_already_renewed_token(self):
update_config(str(self.path), {'access_token': 'new-token'})
with mock.patch.object(skill, '_auto_device_login_unlocked') as login:
cfg = skill.get_config()
self.assertTrue(skill._auto_device_login(cfg, rejected_token='old-token'))
login.assert_not_called()
self.assertEqual(cfg['access_token'], 'new-token')
def test_revoked_device_returns_failure(self):
with mock.patch.object(skill, '_auto_device_login_unlocked', return_value=False), mock.patch.object(skill.requests, 'post', return_value=response(401)) as post:
self.assertFalse(skill.query('HENLO', 'execute_query')['success'])
self.assertEqual(post.call_count, 1)
def test_real_device_challenge_signing_preserves_default_customer(self):
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
from cryptography.hazmat.primitives.serialization import Encoding, PrivateFormat, NoEncryption
private = Ed25519PrivateKey.generate()
key = {'device_id': 'test-device', 'key_fingerprint': 'test-fingerprint',
'private_key': base64.b64encode(private.private_bytes(Encoding.Raw, PrivateFormat.Raw, NoEncryption())).decode('ascii')}
self.prepare('query', '--client', 'WEIRUI', 'SELECT 1 FROM DUAL')
update_config(str(self.path), {'expires_at': '2000-01-01T00:00:00+00:00'})
responses = [response(payload={'success': True, 'challenge': 'test-challenge'}),
response(payload={'success': True, 'access_token': 'signed-token', 'expires_at': '2999-01-01T00:00:00+00:00', 'current_client': {'code': 'WEIRUI'}})]
with mock.patch.object(skill, '_load_device_key', return_value=key), mock.patch.object(skill.requests, 'post', side_effect=responses) as post:
self.assertTrue(skill.ensure_logged_in())
login = post.call_args.kwargs['json']
private.public_key().verify(base64.b64decode(login['signature']), b'test-challenge')
self.assertEqual(login['client_code'], 'WEIRUI')
cfg = read_config(str(self.path))
self.assertEqual(cfg['access_token'], 'signed-token')
self.assertEqual(cfg['client_code'], 'HENLO')
def test_device_denial_is_diagnostic_and_does_not_fallback(self):
from cryptography.hazmat.primitives.asymmetric.ed25519 import Ed25519PrivateKey
from cryptography.hazmat.primitives.serialization import Encoding, PrivateFormat, NoEncryption
private = Ed25519PrivateKey.generate()
key = {'device_id': 'test-device', 'key_fingerprint': 'test-fingerprint',
'private_key': base64.b64encode(private.private_bytes(Encoding.Raw, PrivateFormat.Raw, NoEncryption())).decode('ascii')}
cfg = skill.get_config()
stderr = io.StringIO()
with mock.patch.object(skill, '_load_device_key', return_value=key), mock.patch.object(skill.requests, 'post', return_value=response(403, {'success': False, 'error': 'device revoked'})) as post, redirect_stderr(stderr):
self.assertFalse(skill._auto_device_login(cfg, rejected_token='old-token'))
self.assertEqual(post.call_count, 1)
self.assertIn('device revoked', stderr.getvalue())
self.assertEqual(cfg['access_token'], 'old-token')
def test_management_transport_failure_is_not_replayed(self):
cfg = skill.get_config()
cfg['upgrade_admin_token'] = 'test-admin-token'
skill.save_config(cfg)
with mock.patch.object(skill.requests, 'post', side_effect=skill.requests.Timeout) as post:
self.assertFalse(skill.agent_update('HENLO')['success'])
self.assertEqual(post.call_count, 1)
def test_admin_401_does_not_renew_user_session(self):
cfg = skill.get_config()
cfg['upgrade_admin_token'] = 'test-admin-token'
skill.save_config(cfg)
with mock.patch.object(skill.requests, 'post', return_value=response(401)) as post, mock.patch.object(skill, '_auto_device_login') as login:
self.assertFalse(skill.agent_update('HENLO')['success'])
self.assertEqual(post.call_count, 1)
login.assert_not_called()
def test_timeout_not_replayed_and_connection_failures_are_bounded(self):
with mock.patch.object(skill.requests, 'post', side_effect=skill.requests.Timeout) as post, mock.patch.object(skill.time, 'sleep') as sleep:
self.assertFalse(skill.query('HENLO', 'execute_query')['success'])
self.assertEqual(post.call_count, 1)
sleep.assert_not_called()
with mock.patch.object(skill.requests, 'post', side_effect=skill.requests.ConnectionError) as post, mock.patch.object(skill.time, 'sleep'):
self.assertFalse(skill.query('HENLO', 'execute_query', max_retries=1)['success'])
self.assertEqual(post.call_count, 2)
def test_json_query_with_renewal_has_clean_stdout(self):
stdout, stderr = io.StringIO(), io.StringIO()
def renew(cfg):
print('renew progress', file=sys.stderr)
cfg['access_token'] = 'new-token'
skill.save_config(cfg)
return True
argv = ['skill', 'query', '--client', 'WEIRUI', '--json', 'SELECT 1 FROM DUAL']
with mock.patch.object(sys, 'argv', argv), mock.patch.object(skill.requests, 'post', side_effect=[response(401), response()]), mock.patch.object(skill, '_auto_device_login_unlocked', side_effect=renew), redirect_stdout(stdout), redirect_stderr(stderr):
skill.main()
self.assertEqual(json.loads(stdout.getvalue())['client_code'], 'WEIRUI')
self.assertIn('renew progress', stderr.getvalue())
def test_status_and_refresh_use_one_clients_request(self):
with mock.patch.object(skill.requests, 'get', return_value=response(payload={'success': True, 'clients': [{'code': 'HENLO', 'online': True}]})) as get, redirect_stdout(io.StringIO()):
skill.cmd_status()
self.assertEqual(get.call_count, 1)
self.assertTrue(get.call_args.args[0].endswith('/api/clients'))
def test_brief_capabilities_preserve_full_contract(self):
results = []
for brief in (False, True):
output = io.StringIO()
with redirect_stdout(output):
skill.cmd_capabilities(type('Args', (), {'json': True, 'brief': brief})())
results.append(output.getvalue())
full, brief = map(json.loads, results)
self.assertEqual([c['name'] for c in full['commands']], [c['name'] for c in brief['commands']])
self.assertIn('parameters', full['commands'][0])
self.assertNotIn('parameters', brief['commands'][0])
self.assertLess(len(results[1]), len(results[0]) // 2)
def test_stale_snapshots_merge_only_changed_fields(self):
a, b = skill.get_config(), skill.get_config()
a['access_token'] = 'new-token'
b['client_code'] = b['server_id'] = 'WEIRUI'
skill.save_config(a)
skill.save_config(b)
result = read_config(str(self.path))
self.assertEqual(result['access_token'], 'new-token')
self.assertEqual(result['client_code'], 'WEIRUI')
def test_initialization_does_not_replace_existing_credentials(self):
cfg = update_config(str(self.path), {'access_token': '', 'client_code': ''}, initialize=True)
self.assertEqual(cfg['access_token'], 'old-token')
self.assertEqual(read_config(str(self.path))['client_code'], 'HENLO')
def test_template_initialization_keeps_all_fields(self):
self.path.unlink()
(self.directory / 'config.template.json').write_text(json.dumps({'transit_url': 'https://template.example', 'server_id': '', 'custom': 'keep'}))
with mock.patch.object(skill, 'read_config', wraps=read_config) as read:
# The CWD fallback may exist in the repository; make only that path absent.
def read_without_cwd(path):
if path == 'config.json':
raise FileNotFoundError(path)
return read_config(path)
read.side_effect = read_without_cwd
cfg = skill.get_config()
self.assertEqual(cfg['custom'], 'keep')
self.assertEqual(read_config(str(self.path))['transit_url'], 'https://template.example')
def test_streaming_get_401_keeps_download_options(self):
old, new = response(401), response()
def renew(cfg):
cfg['access_token'] = 'new-token'
skill.save_config(cfg)
return True
with mock.patch.object(skill.requests, 'get', side_effect=[old, new]) as get, mock.patch.object(skill, '_auto_device_login_unlocked', side_effect=renew):
result = skill._http_get('https://transit.example/api/awr/download', params={'client_code': 'WEIRUI', 'date': '20261008'}, headers=skill.make_headers(), stream=True, timeout=120)
self.assertIs(result, new)
old.close.assert_called_once()
self.assertTrue(get.call_args.kwargs['stream'])
self.assertEqual(get.call_args.kwargs['params']['client_code'], 'WEIRUI')
def test_lock_wait_is_bounded(self):
path = str(self.directory / 'test.lock')
with file_lock(path):
with self.assertRaises(TimeoutError):
with file_lock(path, timeout=.1):
self.fail('lock acquired twice')
def start_workers(self, mode):
processes = [subprocess.Popen([sys.executable, __file__, '--worker', mode, str(self.directory), target], stdout=subprocess.PIPE, stderr=subprocess.PIPE) for target in ('HENLO', 'WEIRUI')]
for process in processes:
self.addCleanup(lambda p=process: p.kill() if p.poll() is None else None)
deadline = time.monotonic() + 15
while not all((self.directory / (target + '.ready')).exists() for target in ('HENLO', 'WEIRUI')):
if time.monotonic() > deadline or any(p.poll() is not None for p in processes):
self.fail('workers failed to start')
time.sleep(.01)
(self.directory / 'go').touch()
return processes
def assert_workers(self, processes):
for process in processes:
stdout, stderr = process.communicate(timeout=15)
self.assertEqual(process.returncode, 0, stderr.decode('utf-8', errors='replace'))
def test_two_processes_share_one_expired_login_and_keep_targets(self):
update_config(str(self.path), {'expires_at': '2000-01-01T00:00:00+00:00'})
processes = self.start_workers('renew')
update_config(str(self.path), {'client_code': 'OTHER', 'server_id': 'OTHER'})
self.assert_workers(processes)
self.assertEqual((self.directory / 'renewals').read_text().splitlines(), ['renew'])
self.assertEqual(read_config(str(self.path))['client_code'], 'OTHER')
def test_two_processes_share_one_401_renewal(self):
self.assert_workers(self.start_workers('renew'))
self.assertEqual((self.directory / 'renewals').read_text().splitlines(), ['renew'])
def test_two_processes_merge_configuration_updates(self):
self.assert_workers(self.start_workers('merge'))
cfg = read_config(str(self.path))
self.assertEqual(cfg['HENLO'], 'HENLO')
self.assertEqual(cfg['WEIRUI'], 'WEIRUI')
def test_crashed_process_releases_lock(self):
processes = self.start_workers('crash')
deadline = time.monotonic() + 10
while not (self.directory / 'locked').exists():
if time.monotonic() > deadline:
self.fail('lock not acquired')
time.sleep(.01)
for process in processes:
process.kill()
process.communicate(timeout=5)
with file_lock(str(self.directory / 'crash.lock'), timeout=.5):
pass
if __name__ == '__main__':
if len(sys.argv) > 1 and sys.argv[1] == '--worker':
worker(*sys.argv[2:])
else:
unittest.main()