aboutsummaryrefslogtreecommitdiffstats
path: root/src/test_utils.py
blob: 1db23b2e1bf9378e63b99d8cdc5ca3d1b9d82a77 (plain)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
import random
from dataclasses import dataclass
import tempfile
import sys
import os

def make_test_file(name: str, table: str) -> str:
  return f"""<?xml version="1.0" encoding="utf-8"?>
<circuit>
  <version>2</version>
  <attributes/>
  <visualElements>
    <visualElement>
      <elementName>Testcase</elementName>
      <elementAttributes>
        <entry>
          <string>Label</string>
          <string>{name}</string>
        </entry>
        <entry>
          <string>Testdata</string>
          <testData>
            <dataString>{table}</dataString>
          </testData>
        </entry>
      </elementAttributes>
      <pos x="200" y="200"/>
    </visualElement>
  </visualElements>
  <wires/>
  <measurementOrdering/>
</circuit>
  """

@dataclass
class Param:
  name: str
  min: int
  max: int

class TestGroup:
  def __init__(self, inputs: list[Param], outputs: list[str]):
    self.inputs = inputs
    self.outputs = outputs
    self.cases = []

  def generate_cases(self, n = 1000):
    for _ in range(n):
      result = {}
      for param in self.inputs:
        val = random.randrange(param.min, param.max + 1)
        result[param.name] = val

      self.cases.append(result)

  def process(self):
    raise NotImplemented()

  def to_table(self) -> str:
    names: list[str] = []
    for param in self.inputs:
      names.append(param.name)
    for output in self.outputs:
      names.append(output)

    ret = " ".join(names)
    ret += "\n"

    for case in self.cases:
      row_vals = []
      for param_name in names:
        row_vals.append(str(case[param_name]))

      ret += " ".join(row_vals)
      ret += "\n"

    return ret

def run_test(name: str, table: str):
  if len(sys.argv) != 2:
    print(f"usage: {sys.argv[0]} <circuit.dig>")
    sys.exit(1)

  circuit_path = sys.argv[1]
  _, test_path = tempfile.mkstemp(".dig")

  fc = make_test_file(name, table)
  with open(test_path, "w") as f:
    f.write(fc)

  stream = os.popen(f"make run-single-test circ={circuit_path} tests={test_path}")
  output = stream.read()
  print(output)