4949from cms import utf8_decoder
5050from 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
5454from cms .db .filecacher import FileCacher
5555from cmscommon .archive import Archive
5656from 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
0 commit comments