diff --git a/virtualsmartcard/src/vpicc/virtualsmartcard/CardGenerator.py b/virtualsmartcard/src/vpicc/virtualsmartcard/CardGenerator.py index 5ff1841..b810af5 100644 --- a/virtualsmartcard/src/vpicc/virtualsmartcard/CardGenerator.py +++ b/virtualsmartcard/src/vpicc/virtualsmartcard/CardGenerator.py @@ -382,16 +382,19 @@ class CardGenerator(object): def loadCard(self, filename): """Load a card from disk""" - db = anydbm.open(filename, 'r') - if self.password is None: self.password = getpass.getpass("Please enter your password:") - serializedMF = read_protected_string(db["mf"], self.password) - serializedSAM = read_protected_string(db["sam"], self.password) + db = anydbm.open(filename, 'r') + try: + serializedMF = read_protected_string(db["mf"], self.password) + serializedSAM = read_protected_string(db["sam"], self.password) + self.type = db["type"] + finally: + db.close() + self.sam = loads(serializedSAM) self.mf = loads(serializedMF) - self.type = db["type"] def saveCard(self, filename): """Save the currently running card to disk""" @@ -413,11 +416,13 @@ class CardGenerator(object): protectedSAM = protect_string(sam_string, self.password) db = anydbm.open(filename, 'c') - db["mf"] = protectedMF - db["sam"] = protectedSAM - db["type"] = self.type - db["version"] = "0.1" - db.close() + try: + db["mf"] = protectedMF + db["sam"] = protectedSAM + db["type"] = self.type + db["version"] = "0.1" + finally: + db.close() if __name__ == "__main__": from optparse import OptionParser diff --git a/virtualsmartcard/src/vpicc/virtualsmartcard/tests/CardGenerator_test.py b/virtualsmartcard/src/vpicc/virtualsmartcard/tests/CardGenerator_test.py index b9bf39f..ecb8be2 100644 --- a/virtualsmartcard/src/vpicc/virtualsmartcard/tests/CardGenerator_test.py +++ b/virtualsmartcard/src/vpicc/virtualsmartcard/tests/CardGenerator_test.py @@ -17,9 +17,10 @@ # virtualsmartcard. If not, see . # -import unittest -import tempfile +import anydbm import os +import tempfile +import unittest from virtualsmartcard.CardGenerator import CardGenerator class TestNPACardGenerator(unittest.TestCase): @@ -29,15 +30,13 @@ class TestNPACardGenerator(unittest.TestCase): self.nPA_generator = CardGenerator('nPA') self.nPA_generator.password = "TestPassword" - def tearDown(self): - os.unlink(self.filename) - def test_nPA_creation(self): self.nPA_generator.generateCard() self.nPA_generator.saveCard(self.filename) mf, sam = self.nPA_generator.getCard() self.assertIsNotNone(mf) self.assertIsNotNone(sam) + os.unlink(self.filename) def test_load_nPA_from_file_nPA_from_file(self): self.nPA_generator.generateCard() @@ -48,6 +47,11 @@ class TestNPACardGenerator(unittest.TestCase): mf, sam = local_generator.getCard() self.assertIsNotNone(mf) self.assertIsNotNone(sam) + os.unlink(self.filename) + + def test_load_nonexistent_file(self): + with self.assertRaises(anydbm.error): + self.nPA_generator.loadCard(self.filename) if __name__ == '__main__': unittest.main()