249 lines
8.2 KiB
Python
249 lines
8.2 KiB
Python
import json
|
|
|
|
from analytics import AnalyticsRepository
|
|
from redis_task_distribute import RedisTaskDispatcher
|
|
|
|
|
|
class FakeRedis:
|
|
def __init__(self):
|
|
self.lists = {}
|
|
self.hashes = {}
|
|
self.sets = {}
|
|
|
|
def delete(self, key):
|
|
self.lists.pop(key, None)
|
|
self.hashes.pop(key, None)
|
|
self.sets.pop(key, None)
|
|
|
|
def lpush(self, key, value):
|
|
self.lists.setdefault(key, []).insert(0, value)
|
|
|
|
def rpush(self, key, value):
|
|
self.lists.setdefault(key, []).append(value)
|
|
|
|
def rpop(self, key):
|
|
values = self.lists.setdefault(key, [])
|
|
return values.pop() if values else None
|
|
|
|
def lrange(self, key, start, end):
|
|
values = self.lists.get(key, [])
|
|
stop = None if end == -1 else end + 1
|
|
return values[start:stop]
|
|
|
|
def llen(self, key):
|
|
return len(self.lists.get(key, []))
|
|
|
|
def lpos(self, key, value):
|
|
try:
|
|
return self.lists.get(key, []).index(value)
|
|
except ValueError:
|
|
return None
|
|
|
|
def lrem(self, key, count, value):
|
|
values = self.lists.get(key, [])
|
|
original_len = len(values)
|
|
self.lists[key] = [item for item in values if item != value]
|
|
return original_len - len(self.lists[key])
|
|
|
|
def hset(self, key, field, value):
|
|
self.hashes.setdefault(key, {})[field] = value
|
|
|
|
def hget(self, key, field):
|
|
return self.hashes.get(key, {}).get(field)
|
|
|
|
def hgetall(self, key):
|
|
return dict(self.hashes.get(key, {}))
|
|
|
|
def hvals(self, key):
|
|
return list(self.hashes.get(key, {}).values())
|
|
|
|
def hlen(self, key):
|
|
return len(self.hashes.get(key, {}))
|
|
|
|
def hexists(self, key, field):
|
|
return field in self.hashes.get(key, {})
|
|
|
|
def hdel(self, key, field):
|
|
self.hashes.get(key, {}).pop(field, None)
|
|
|
|
def sadd(self, key, value):
|
|
self.sets.setdefault(key, set()).add(value)
|
|
|
|
def srem(self, key, value):
|
|
self.sets.setdefault(key, set()).discard(value)
|
|
|
|
def scard(self, key):
|
|
return len(self.sets.get(key, set()))
|
|
|
|
def smembers(self, key):
|
|
return set(self.sets.get(key, set()))
|
|
|
|
|
|
class StubAnalytics:
|
|
def __init__(self, pending=None, model=None):
|
|
self.pending = pending or []
|
|
self.model = model or []
|
|
|
|
def list_pending_collection_tasks(self):
|
|
return list(self.pending)
|
|
|
|
def list_model_eligible_apps(self):
|
|
return list(self.model)
|
|
|
|
|
|
def make_dispatcher(redis_conn=None, analytics=None):
|
|
dispatcher = RedisTaskDispatcher.__new__(RedisTaskDispatcher)
|
|
dispatcher.redis = redis_conn or FakeRedis()
|
|
dispatcher.analytics = analytics or StubAnalytics()
|
|
dispatcher.worker_inventory = {}
|
|
dispatcher.managed_worker_ids = set()
|
|
dispatcher.worker_online_timeout = 300
|
|
dispatcher.task_routing_rules = {"package_name": {}, "task_key": {}}
|
|
return dispatcher
|
|
|
|
|
|
def pending_entry(app_name, package_name, task_queue="default"):
|
|
return {
|
|
"app_name": app_name,
|
|
"package_name": package_name,
|
|
"task_queue": task_queue,
|
|
"task_payload": {
|
|
"app_name": app_name,
|
|
"package_name": package_name,
|
|
"country_code": "US",
|
|
},
|
|
}
|
|
|
|
|
|
def test_statistics_and_dashboard_read_three_priority_queues():
|
|
redis_conn = FakeRedis()
|
|
dispatcher = make_dispatcher(redis_conn=redis_conn)
|
|
for app_name, package_name, queue_name in (
|
|
("High App", "com.example.high", "task:queue:high"),
|
|
("Default App", "com.example.default", "task:queue:default"),
|
|
("Low App", "com.example.low", "task:queue:low"),
|
|
):
|
|
task_key = f"{app_name}_{package_name}"
|
|
redis_conn.hset("task:details", task_key, json.dumps({
|
|
"app_name": app_name,
|
|
"package_name": package_name,
|
|
"task_queue": queue_name.rsplit(":", 1)[-1],
|
|
}))
|
|
redis_conn.lpush(queue_name, task_key)
|
|
|
|
stats = dispatcher.get_statistics()
|
|
dashboard_tasks = dispatcher.get_dashboard_tasks()
|
|
|
|
assert stats["pending"] == 3
|
|
assert [item["task_key"] for item in dashboard_tasks["pending"]] == [
|
|
"High App_com.example.high",
|
|
"Default App_com.example.default",
|
|
"Low App_com.example.low",
|
|
]
|
|
|
|
|
|
def test_refresh_requeues_existing_completed_high_priority_task():
|
|
entry = pending_entry("Hot App", "com.example.hot", task_queue="high")
|
|
redis_conn = FakeRedis()
|
|
dispatcher = make_dispatcher(redis_conn=redis_conn, analytics=StubAnalytics(pending=[entry]))
|
|
task_key = "Hot App_com.example.hot"
|
|
redis_conn.hset("task:details", task_key, json.dumps({"app_name": "Old", "package_name": "com.example.hot"}))
|
|
redis_conn.hset("task:status", task_key, json.dumps({"status": "completed", "retry_count": 2}))
|
|
redis_conn.sadd("task:completed", task_key)
|
|
|
|
assert dispatcher.refresh_tasks_from_app_summary() == 1
|
|
|
|
assert redis_conn.lrange("task:queue:high", 0, -1) == [task_key]
|
|
assert redis_conn.scard("task:completed") == 0
|
|
assert json.loads(redis_conn.hget("task:status", task_key))["status"] == "pending"
|
|
task_details = json.loads(redis_conn.hget("task:details", task_key))
|
|
assert task_details["app_name"] == "Hot App"
|
|
assert task_details["task_queue"] == "high"
|
|
|
|
|
|
def test_refresh_does_not_duplicate_running_task():
|
|
entry = pending_entry("Running App", "com.example.running", task_queue="high")
|
|
redis_conn = FakeRedis()
|
|
dispatcher = make_dispatcher(redis_conn=redis_conn, analytics=StubAnalytics(pending=[entry]))
|
|
task_key = "Running App_com.example.running"
|
|
redis_conn.hset("task:status", task_key, json.dumps({"status": "running", "worker_id": "worker-1"}))
|
|
redis_conn.hset("worker:tasks", "worker-1", task_key)
|
|
|
|
assert dispatcher.refresh_tasks_from_app_summary() == 0
|
|
assert redis_conn.lrange("task:queue:high", 0, -1) == []
|
|
|
|
|
|
def test_direct_mode_payload_does_not_default_to_local_source(monkeypatch):
|
|
import redis_task_distribute as dispatcher_module
|
|
|
|
monkeypatch.setattr(dispatcher_module, "APK_DOWNLOAD_MODE", "direct")
|
|
dispatcher = make_dispatcher()
|
|
dispatcher._get_apk_registry = lambda: (_ for _ in ()).throw(AssertionError("registry should not be used in direct mode"))
|
|
|
|
payload = dispatcher._task_to_payload(
|
|
"Direct App_com.example.direct",
|
|
{
|
|
"app_name": "Direct App",
|
|
"package_name": "com.example.direct",
|
|
"country_code": "US",
|
|
},
|
|
)
|
|
|
|
assert payload["available_sources"] == ["google_play", "apkpure"]
|
|
assert payload["local_apk_dir"] == ""
|
|
|
|
|
|
def test_minio_mode_payload_can_use_local_source(monkeypatch):
|
|
import redis_task_distribute as dispatcher_module
|
|
|
|
monkeypatch.setattr(dispatcher_module, "APK_DOWNLOAD_MODE", "minio")
|
|
dispatcher = make_dispatcher()
|
|
|
|
class _Registry:
|
|
@staticmethod
|
|
def is_fresh_enough(package_name, last_updated):
|
|
return False
|
|
|
|
dispatcher._get_apk_registry = lambda: _Registry()
|
|
|
|
payload = dispatcher._task_to_payload(
|
|
"Cached App_com.example.cached",
|
|
{
|
|
"app_name": "Cached App",
|
|
"package_name": "com.example.cached",
|
|
"country_code": "US",
|
|
},
|
|
)
|
|
|
|
assert payload["available_sources"] == ["google_play", "local", "apkpure"]
|
|
|
|
|
|
def test_direct_mode_download_error_does_not_mark_pending_apk(monkeypatch):
|
|
import redis_task_distribute as dispatcher_module
|
|
|
|
monkeypatch.setattr(dispatcher_module, "APK_DOWNLOAD_MODE", "direct")
|
|
dispatcher = make_dispatcher()
|
|
dispatcher._get_apk_registry = lambda: (_ for _ in ()).throw(AssertionError("registry should not be used in direct mode"))
|
|
|
|
dispatcher._mark_download_error_awaiting_apk("com.example.direct")
|
|
|
|
|
|
def test_analytics_jobs_accept_scheduled_after(tmp_path):
|
|
repo = AnalyticsRepository(db_path=str(tmp_path / "analytics.sqlite3"))
|
|
scheduled_after = 12345.0
|
|
|
|
job = repo.create_job(
|
|
job_type="incremental",
|
|
status="queued",
|
|
package_name="com.example.active",
|
|
scheduled_after=scheduled_after,
|
|
)
|
|
|
|
with repo._connect() as connection:
|
|
row = connection.execute(
|
|
"SELECT scheduled_after FROM analytics_job WHERE id = ?",
|
|
(job["id"],),
|
|
).fetchone()
|
|
|
|
assert row["scheduled_after"] == scheduled_after
|