Merge pull request #35 from d0/master

Minor improvements to CardGenerator and the corresponding unit test
This commit is contained in:
Frank Morgner
2015-01-08 12:12:02 +01:00
2 changed files with 65 additions and 58 deletions

View File

@@ -380,19 +380,21 @@ class CardGenerator(object):
if sam != None: if sam != None:
self.sam = sam self.sam = sam
def loadCard(self, filename): def loadCard(self, filename):
"""Load a card from disk""" """Load a card from disk"""
db = anydbm.open(filename, 'r')
if self.password is None: if self.password is None:
self.password = getpass.getpass("Please enter your password:") self.password = getpass.getpass("Please enter your password:")
db = anydbm.open(filename, 'r')
try:
serializedMF = read_protected_string(db["mf"], self.password) serializedMF = read_protected_string(db["mf"], self.password)
serializedSAM = read_protected_string(db["sam"], self.password) serializedSAM = read_protected_string(db["sam"], self.password)
self.type = db["type"]
finally:
db.close()
self.sam = loads(serializedSAM) self.sam = loads(serializedSAM)
self.mf = loads(serializedMF) self.mf = loads(serializedMF)
self.type = db["type"]
def saveCard(self, filename): def saveCard(self, filename):
"""Save the currently running card to disk""" """Save the currently running card to disk"""
@@ -414,10 +416,12 @@ class CardGenerator(object):
protectedSAM = protect_string(sam_string, self.password) protectedSAM = protect_string(sam_string, self.password)
db = anydbm.open(filename, 'c') db = anydbm.open(filename, 'c')
try:
db["mf"] = protectedMF db["mf"] = protectedMF
db["sam"] = protectedSAM db["sam"] = protectedSAM
db["type"] = self.type db["type"] = self.type
db["version"] = "0.1" db["version"] = "0.1"
finally:
db.close() db.close()
def readDatagroups(self, datasetfile): def readDatagroups(self, datasetfile):
@@ -430,7 +434,6 @@ class CardGenerator(object):
splitLine = line.split("=") splitLine = line.split("=")
# we don't want to have the newline char from dataset file as part of the value!! # we don't want to have the newline char from dataset file as part of the value!!
self.datagroups[splitLine[0]] = splitLine[1].rstrip("\n\r") self.datagroups[splitLine[0]] = splitLine[1].rstrip("\n\r")
f.close()
if __name__ == "__main__": if __name__ == "__main__":
from optparse import OptionParser from optparse import OptionParser

View File

@@ -17,9 +17,10 @@
# virtualsmartcard. If not, see <http://www.gnu.org/licenses/>. # virtualsmartcard. If not, see <http://www.gnu.org/licenses/>.
# #
import unittest import anydbm
import tempfile
import os import os
import tempfile
import unittest
from virtualsmartcard.CardGenerator import CardGenerator from virtualsmartcard.CardGenerator import CardGenerator
class TestNPACardGenerator(unittest.TestCase): class TestNPACardGenerator(unittest.TestCase):
@@ -29,15 +30,13 @@ class TestNPACardGenerator(unittest.TestCase):
self.nPA_generator = CardGenerator('nPA') self.nPA_generator = CardGenerator('nPA')
self.nPA_generator.password = "TestPassword" self.nPA_generator.password = "TestPassword"
def tearDown(self):
os.unlink(self.filename)
def test_nPA_creation(self): def test_nPA_creation(self):
self.nPA_generator.generateCard() self.nPA_generator.generateCard()
self.nPA_generator.saveCard(self.filename) self.nPA_generator.saveCard(self.filename)
mf, sam = self.nPA_generator.getCard() mf, sam = self.nPA_generator.getCard()
self.assertIsNotNone(mf) self.assertIsNotNone(mf)
self.assertIsNotNone(sam) self.assertIsNotNone(sam)
os.unlink(self.filename)
def test_load_nPA_from_file_nPA_from_file(self): def test_load_nPA_from_file_nPA_from_file(self):
self.nPA_generator.generateCard() self.nPA_generator.generateCard()
@@ -48,6 +47,11 @@ class TestNPACardGenerator(unittest.TestCase):
mf, sam = local_generator.getCard() mf, sam = local_generator.getCard()
self.assertIsNotNone(mf) self.assertIsNotNone(mf)
self.assertIsNotNone(sam) 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__': if __name__ == '__main__':
unittest.main() unittest.main()