Skip to content

Commit 99f2f2b

Browse files
ranaldmiaoandreyv
authored andcommitted
Add no-users option to DumpImporter and DumpExporter (cms-dev#1165)
* feat: Allow DumpExporter to only export tasks * fix: Dump exporter tests * fix: Add DumpExporterTest for skip_users, fix bug * feat: Add no-users option to DumpImporter * fixup Co-authored-by: Andrey Vihrov <andrey.vihrov@gmail.com>
1 parent 5e0cbd0 commit 99f2f2b

5 files changed

Lines changed: 106 additions & 24 deletions

File tree

cms/db/util.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -274,8 +274,8 @@ def get_datasets_to_judge(task):
274274

275275
def enumerate_files(
276276
session, contest=None,
277-
skip_submissions=False, skip_user_tests=False, skip_print_jobs=False,
278-
skip_generated=False):
277+
skip_submissions=False, skip_user_tests=False, skip_users=False,
278+
skip_print_jobs=False, skip_generated=False):
279279
"""Enumerate all the files (by digest) referenced by the
280280
contest.
281281
@@ -302,7 +302,7 @@ def enumerate_files(
302302
queries.append(dataset_q.join(Dataset.testcases)
303303
.with_entities(Testcase.output))
304304

305-
if not skip_submissions:
305+
if not skip_submissions and not skip_users:
306306
submission_q = task_q.join(Task.submissions)
307307
queries.append(submission_q.join(Submission.files)
308308
.with_entities(File.digest))
@@ -312,7 +312,7 @@ def enumerate_files(
312312
.join(SubmissionResult.executables)
313313
.with_entities(Executable.digest))
314314

315-
if not skip_user_tests:
315+
if not skip_user_tests and not skip_users:
316316
user_test_q = task_q.join(Task.user_tests)
317317
queries.append(user_test_q.with_entities(UserTest.input))
318318
queries.append(user_test_q.join(UserTest.files)
@@ -328,7 +328,7 @@ def enumerate_files(
328328
.filter(UserTestResult.output != None)
329329
.with_entities(UserTestResult.output))
330330

331-
if not skip_print_jobs:
331+
if not skip_print_jobs and not skip_users:
332332
queries.append(contest_q.join(Contest.participations)
333333
.join(Participation.printjobs)
334334
.with_entities(PrintJob.digest))

cmscontrib/DumpExporter.py

Lines changed: 23 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -47,7 +47,7 @@
4747
from cms.db import version as model_version, Codename, Filename, \
4848
FilenameSchema, FilenameSchemaArray, Digest, SessionGen, Contest, User, \
4949
Task, Submission, UserTest, SubmissionResult, UserTestResult, PrintJob, \
50-
enumerate_files
50+
Announcement, Participation, enumerate_files
5151
from cms.db.filecacher import FileCacher
5252
from cmscommon.datetime import make_timestamp
5353
from cmscommon.digest import path_digest
@@ -136,13 +136,16 @@ class DumpExporter:
136136

137137
def __init__(self, contest_ids, export_target,
138138
dump_files, dump_model, skip_generated,
139-
skip_submissions, skip_user_tests, skip_print_jobs):
139+
skip_submissions, skip_user_tests, skip_users, skip_print_jobs):
140140
if contest_ids is None:
141141
with SessionGen() as session:
142142
contests = session.query(Contest).all()
143143
self.contests_ids = [contest.id for contest in contests]
144-
users = session.query(User).all()
145-
self.users_ids = [user.id for user in users]
144+
if not skip_users:
145+
users = session.query(User).all()
146+
self.users_ids = [user.id for user in users]
147+
else:
148+
self.users_ids = []
146149
tasks = session.query(Task)\
147150
.filter(Task.contest_id.is_(None)).all()
148151
self.tasks_ids = [task.id for task in tasks]
@@ -158,6 +161,7 @@ def __init__(self, contest_ids, export_target,
158161
self.skip_generated = skip_generated
159162
self.skip_submissions = skip_submissions
160163
self.skip_user_tests = skip_user_tests
164+
self.skip_users = skip_users
161165
self.skip_print_jobs = skip_print_jobs
162166
self.export_target = export_target
163167

@@ -208,6 +212,7 @@ def do_export(self):
208212
session, contest,
209213
skip_submissions=self.skip_submissions,
210214
skip_user_tests=self.skip_user_tests,
215+
skip_users=self.skip_users,
211216
skip_print_jobs=self.skip_print_jobs,
212217
skip_generated=self.skip_generated)
213218
for file_ in files:
@@ -317,6 +322,17 @@ class of the given object), an item for each column property
317322
if self.skip_user_tests and other_cls is UserTest:
318323
continue
319324

325+
if self.skip_users:
326+
skip = False
327+
# User-related classes reachable from root
328+
for rel_class in [Participation, Submission, UserTest,
329+
Announcement]:
330+
if other_cls is rel_class:
331+
skip = True
332+
break
333+
if skip:
334+
continue
335+
320336
# Skip print jobs if requested
321337
if self.skip_print_jobs and other_cls is PrintJob:
322338
continue
@@ -397,6 +413,8 @@ def main():
397413
help="don't export submissions")
398414
parser.add_argument("-U", "--no-user-tests", action="store_true",
399415
help="don't export user tests")
416+
parser.add_argument("-X", "--no-users", action="store_true",
417+
help="don't export users")
400418
parser.add_argument("-P", "--no-print-jobs", action="store_true",
401419
help="don't export print jobs")
402420
parser.add_argument("export_target", action="store",
@@ -412,6 +430,7 @@ def main():
412430
skip_generated=args.no_generated,
413431
skip_submissions=args.no_submissions,
414432
skip_user_tests=args.no_user_tests,
433+
skip_users=args.no_users,
415434
skip_print_jobs=args.no_print_jobs)
416435
success = exporter.do_export()
417436
return 0 if success is True else 1

cmscontrib/DumpImporter.py

Lines changed: 29 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -49,8 +49,8 @@
4949
from cms import utf8_decoder
5050
from cms.db import version as model_version, Codename, Filename, \
5151
FilenameSchema, FilenameSchemaArray, Digest, SessionGen, Contest, \
52-
Submission, SubmissionResult, UserTest, UserTestResult, PrintJob, init_db, \
53-
drop_db, enumerate_files
52+
Submission, SubmissionResult, User, Participation, UserTest, \
53+
UserTestResult, PrintJob, Announcement, init_db, drop_db, enumerate_files
5454
from cms.db.filecacher import FileCacher
5555
from cmscommon.archive import Archive
5656
from cmscommon.datetime import make_datetime
@@ -128,13 +128,14 @@ class DumpImporter:
128128

129129
def __init__(self, drop, import_source,
130130
load_files, load_model, skip_generated,
131-
skip_submissions, skip_user_tests, skip_print_jobs):
131+
skip_submissions, skip_user_tests, skip_users, skip_print_jobs):
132132
self.drop = drop
133133
self.load_files = load_files
134134
self.load_model = load_model
135135
self.skip_generated = skip_generated
136136
self.skip_submissions = skip_submissions
137137
self.skip_user_tests = skip_user_tests
138+
self.skip_users = skip_users
138139
self.skip_print_jobs = skip_print_jobs
139140

140141
self.import_source = import_source
@@ -233,9 +234,6 @@ def do_import(self):
233234
for id_, data in self.datas.items():
234235
if not id_.startswith("_"):
235236
self.objs[id_] = self.import_object(data)
236-
for id_, data in self.datas.items():
237-
if not id_.startswith("_"):
238-
self.add_relationships(data, self.objs[id_])
239237

240238
for k, v in list(self.objs.items()):
241239

@@ -244,18 +242,28 @@ def do_import(self):
244242
del self.objs[k]
245243

246244
# Skip user_tests if requested
247-
if self.skip_user_tests and isinstance(v, UserTest):
245+
elif self.skip_user_tests and isinstance(v, UserTest):
246+
del self.objs[k]
247+
248+
# Skip users if requested
249+
elif self.skip_users and \
250+
isinstance(v, (User, Participation, Submission,
251+
UserTest, Announcement)):
248252
del self.objs[k]
249253

250254
# Skip print jobs if requested
251-
if self.skip_print_jobs and isinstance(v, PrintJob):
255+
elif self.skip_print_jobs and isinstance(v, PrintJob):
252256
del self.objs[k]
253257

254258
# Skip generated data if requested
255-
if self.skip_generated and \
259+
elif self.skip_generated and \
256260
isinstance(v, (SubmissionResult, UserTestResult)):
257261
del self.objs[k]
258262

263+
for id_, data in self.datas.items():
264+
if not id_.startswith("_") and id_ in self.objs:
265+
self.add_relationships(data, self.objs[id_])
266+
259267
contest_id = list()
260268
contest_files = set()
261269

@@ -266,6 +274,11 @@ def do_import(self):
266274
# that depended on submissions or user tests that we
267275
# might have removed above).
268276
for id_ in self.datas["_objects"]:
277+
278+
# It could have been removed by request
279+
if id_ not in self.objs:
280+
continue
281+
269282
obj = self.objs[id_]
270283
session.add(obj)
271284
session.flush()
@@ -277,6 +290,7 @@ def do_import(self):
277290
skip_submissions=self.skip_submissions,
278291
skip_user_tests=self.skip_user_tests,
279292
skip_print_jobs=self.skip_print_jobs,
293+
skip_users=self.skip_users,
280294
skip_generated=self.skip_generated)
281295

282296
session.commit()
@@ -405,12 +419,12 @@ def add_relationships(self, data, obj):
405419
if val is None:
406420
setattr(obj, prp.key, None)
407421
elif isinstance(val, str):
408-
setattr(obj, prp.key, self.objs[val])
422+
setattr(obj, prp.key, self.objs.get(val))
409423
elif isinstance(val, list):
410-
setattr(obj, prp.key, list(self.objs[i] for i in val))
424+
setattr(obj, prp.key, list(self.objs[i] for i in val if i in self.objs))
411425
elif isinstance(val, dict):
412426
setattr(obj, prp.key,
413-
dict((k, self.objs[v]) for k, v in val.items()))
427+
dict((k, self.objs[v]) for k, v in val.items() if v in self.objs))
414428
else:
415429
raise RuntimeError(
416430
"Unknown RelationshipProperty value: %s" % type(val))
@@ -472,6 +486,8 @@ def main():
472486
help="don't import submissions")
473487
parser.add_argument("-U", "--no-user-tests", action="store_true",
474488
help="don't import user tests")
489+
parser.add_argument("-X", "--no-users", action="store_true",
490+
help="don't import users")
475491
parser.add_argument("-P", "--no-print-jobs", action="store_true",
476492
help="don't import print jobs")
477493
parser.add_argument("import_source", action="store", type=utf8_decoder,
@@ -486,6 +502,7 @@ def main():
486502
skip_generated=args.no_generated,
487503
skip_submissions=args.no_submissions,
488504
skip_user_tests=args.no_user_tests,
505+
skip_users=args.no_users,
489506
skip_print_jobs=args.no_print_jobs)
490507
success = importer.do_import()
491508
return 0 if success is True else 1

cmstestsuite/unit_tests/cmscontrib/DumpExporterTest.py

Lines changed: 21 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -85,7 +85,7 @@ def tearDown(self):
8585
super().tearDown()
8686

8787
def do_export(self, contest_ids, dump_files=True, skip_generated=False,
88-
skip_submissions=False):
88+
skip_submissions=False, skip_users=False):
8989
"""Create an exporter and call do_export in a convenient way"""
9090
r = DumpExporter(
9191
contest_ids,
@@ -95,6 +95,7 @@ def do_export(self, contest_ids, dump_files=True, skip_generated=False,
9595
skip_generated=skip_generated,
9696
skip_submissions=skip_submissions,
9797
skip_user_tests=False,
98+
skip_users=skip_users,
9899
skip_print_jobs=False).do_export()
99100
dump_path = os.path.join(self.target, "contest.json")
100101
try:
@@ -269,6 +270,25 @@ def test_skip_generated(self):
269270
self.assertNotInDump(SubmissionResult)
270271
self.assertFileNotInDump(self.exe_digest)
271272

273+
def test_skip_users(self):
274+
"""Test skipping users.
275+
276+
Should not export users and depending objects.
277+
Should still export contest, tasks and their depending objects.
278+
279+
"""
280+
self.assertTrue(self.do_export(None, skip_users=True))
281+
282+
self.assertInDump(Statement, digest=self.st_digest)
283+
self.assertFileInDump(self.st_digest, self.st_content)
284+
285+
self.assertNotInDump(User)
286+
self.assertNotInDump(Participation)
287+
self.assertNotInDump(Submission)
288+
self.assertNotInDump(SubmissionResult)
289+
self.assertFileNotInDump(self.file_digest)
290+
self.assertFileNotInDump(self.exe_digest)
291+
272292

273293
if __name__ == "__main__":
274294
unittest.main()

cmstestsuite/unit_tests/cmscontrib/DumpImporterTest.py

Lines changed: 28 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,7 @@
2525
# Needs to be first to allow for monkey patching the DB connection string.
2626
from cmstestsuite.unit_tests.databasemixin import DatabaseMixin
2727

28-
from cms.db import Contest, FSObject, Session, version
28+
from cms.db import Contest, User, FSObject, Session, version
2929
from cmscommon.digest import bytes_digest
3030
from cmscontrib.DumpImporter import DumpImporter
3131
from cmstestsuite.unit_tests.filesystemmixin import FileSystemMixin
@@ -136,7 +136,8 @@ def tearDown(self):
136136
super().tearDown()
137137

138138
def do_import(self, drop=False, load_files=True,
139-
skip_generated=False, skip_submissions=False):
139+
skip_generated=False, skip_submissions=False,
140+
skip_users=False):
140141
"""Create an importer and call do_import in a convenient way"""
141142
return DumpImporter(
142143
drop,
@@ -146,6 +147,7 @@ def do_import(self, drop=False, load_files=True,
146147
skip_generated=skip_generated,
147148
skip_submissions=skip_submissions,
148149
skip_user_tests=False,
150+
skip_users=skip_users,
149151
skip_print_jobs=False).do_import()
150152

151153
def write_dump(self, dump):
@@ -195,6 +197,12 @@ def assertContestNotInDb(self, name):
195197
.filter(Contest.name == name).all()
196198
self.assertEqual(len(db_contests), 0)
197199

200+
def assertUserNotInDb(self, username):
201+
"""Assert that the user with the given username is not in the DB."""
202+
db_users = self.session.query(User)\
203+
.filter(User.username == username).all()
204+
self.assertEqual(len(db_users), 0)
205+
198206
def assertFileInDb(self, digest, description, content):
199207
"""Assert that the file with the given data is in the DB."""
200208
fsos = self.session.query(FSObject)\
@@ -279,6 +287,24 @@ def test_import_skip_files(self):
279287
self.assertFileNotInDb(TestDumpImporter.GENERATED_FILE_DIGEST)
280288
self.assertFileNotInDb(TestDumpImporter.NON_GENERATED_FILE_DIGEST)
281289

290+
def test_import_skip_users(self):
291+
"""Test importing everything but not the users."""
292+
self.write_dump(TestDumpImporter.DUMP)
293+
self.write_files(TestDumpImporter.FILES)
294+
295+
self.assertTrue(self.do_import(skip_users=True))
296+
297+
self.assertContestInDb("contestname", "contest description 你好",
298+
[("taskname", "task title")],
299+
[])
300+
self.assertContestInDb(
301+
self.other_contest_name, self.other_contest_description, [], [])
302+
303+
self.assertUserNotInDb("username")
304+
self.assertFileNotInDb(TestDumpImporter.GENERATED_FILE_DIGEST)
305+
self.assertFileNotInDb(TestDumpImporter.NON_GENERATED_FILE_DIGEST)
306+
307+
282308
def test_import_old(self):
283309
"""Test importing an old dump.
284310

0 commit comments

Comments
 (0)