391 lines
20 KiB
Python
391 lines
20 KiB
Python
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()
|