1from __future__ import print_function
2from __future__ import absolute_import
3
4# System modules
5import os
6
7# Third-party modules
8
9# LLDB modules
10import lldb
11from .lldbtest import *
12from . import lldbutil
13from .decorators import *
14
15def source_type(filename):
16    _, extension = os.path.splitext(filename)
17    return {
18        '.c' : 'C_SOURCES',
19        '.cpp' : 'CXX_SOURCES',
20        '.cxx' : 'CXX_SOURCES',
21        '.cc' : 'CXX_SOURCES',
22        '.m' : 'OBJC_SOURCES',
23        '.mm' : 'OBJCXX_SOURCES'
24    }.get(extension, None)
25
26
27class CommandParser:
28    def __init__(self):
29        self.breakpoints = []
30
31    def parse_one_command(self, line):
32        parts = line.split('//%')
33
34        command = None
35        new_breakpoint = True
36
37        if len(parts) == 2:
38            command = parts[1].strip()  # take off whitespace
39            new_breakpoint = parts[0].strip() != ""
40
41        return (command, new_breakpoint)
42
43    def parse_source_files(self, source_files):
44        for source_file in source_files:
45            file_handle = open(source_file)
46            lines = file_handle.readlines()
47            line_number = 0
48            current_breakpoint = None # non-NULL means we're looking through whitespace to find additional commands
49            for line in lines:
50                line_number = line_number + 1 # 1-based, so we do this first
51                (command, new_breakpoint) = self.parse_one_command(line)
52
53                if new_breakpoint:
54                    current_breakpoint = None
55
56                if command != None:
57                    if current_breakpoint == None:
58                        current_breakpoint = {}
59                        current_breakpoint['file_name'] = source_file
60                        current_breakpoint['line_number'] = line_number
61                        current_breakpoint['command'] = command
62                        self.breakpoints.append(current_breakpoint)
63                    else:
64                        current_breakpoint['command'] = current_breakpoint['command'] + "\n" + command
65
66    def set_breakpoints(self, target):
67        for breakpoint in self.breakpoints:
68            breakpoint['breakpoint'] = target.BreakpointCreateByLocation(breakpoint['file_name'], breakpoint['line_number'])
69
70    def handle_breakpoint(self, test, breakpoint_id):
71        for breakpoint in self.breakpoints:
72            if breakpoint['breakpoint'].GetID() == breakpoint_id:
73                test.execute_user_command(breakpoint['command'])
74                return
75
76class InlineTest(TestBase):
77    # Internal implementation
78
79    def getRerunArgs(self):
80        # The -N option says to NOT run a if it matches the option argument, so
81        # if we are using dSYM we say to NOT run dwarf (-N dwarf) and vice versa.
82        if self.using_dsym is None:
83            # The test was skipped altogether.
84            return ""
85        elif self.using_dsym:
86            return "-N dwarf %s" % (self.mydir)
87        else:
88            return "-N dsym %s" % (self.mydir)
89
90    def BuildMakefile(self):
91        if os.path.exists("Makefile"):
92            return
93
94        categories = {}
95
96        for f in os.listdir(os.getcwd()):
97            t = source_type(f)
98            if t:
99                if t in list(categories.keys()):
100                    categories[t].append(f)
101                else:
102                    categories[t] = [f]
103
104        makefile = open("Makefile", 'w+')
105
106        level = os.sep.join([".."] * len(self.mydir.split(os.sep))) + os.sep + "make"
107
108        makefile.write("LEVEL = " + level + "\n")
109
110        for t in list(categories.keys()):
111            line = t + " := " + " ".join(categories[t])
112            makefile.write(line + "\n")
113
114        if ('OBJCXX_SOURCES' in list(categories.keys())) or ('OBJC_SOURCES' in list(categories.keys())):
115            makefile.write("LDFLAGS = $(CFLAGS) -lobjc -framework Foundation\n")
116
117        if ('CXX_SOURCES' in list(categories.keys())):
118            makefile.write("CXXFLAGS += -std=c++11\n")
119
120        makefile.write("include $(LEVEL)/Makefile.rules\n")
121        makefile.write("\ncleanup:\n\trm -f Makefile *.d\n\n")
122        makefile.flush()
123        makefile.close()
124
125    @skipUnlessDarwin
126    def __test_with_dsym(self):
127        self.using_dsym = True
128        self.BuildMakefile()
129        self.buildDsym()
130        self.do_test()
131
132    def __test_with_dwarf(self):
133        self.using_dsym = False
134        self.BuildMakefile()
135        self.buildDwarf()
136        self.do_test()
137
138    def __test_with_dwo(self):
139        self.using_dsym = False
140        self.BuildMakefile()
141        self.buildDwo()
142        self.do_test()
143
144    def execute_user_command(self, __command):
145        exec(__command, globals(), locals())
146
147    def do_test(self):
148        exe_name = "a.out"
149        exe = os.path.join(os.getcwd(), exe_name)
150        source_files = [ f for f in os.listdir(os.getcwd()) if source_type(f) ]
151        target = self.dbg.CreateTarget(exe)
152
153        parser = CommandParser()
154        parser.parse_source_files(source_files)
155        parser.set_breakpoints(target)
156
157        process = target.LaunchSimple(None, None, os.getcwd())
158
159        while lldbutil.get_stopped_thread(process, lldb.eStopReasonBreakpoint):
160            thread = lldbutil.get_stopped_thread(process, lldb.eStopReasonBreakpoint)
161            breakpoint_id = thread.GetStopReasonDataAtIndex (0)
162            parser.handle_breakpoint(self, breakpoint_id)
163            process.Continue()
164
165
166    # Utilities for testcases
167
168    def check_expression (self, expression, expected_result, use_summary = True):
169        value = self.frame().EvaluateExpression (expression)
170        self.assertTrue(value.IsValid(), expression+"returned a valid value")
171        if self.TraceOn():
172            print(value.GetSummary())
173            print(value.GetValue())
174        if use_summary:
175            answer = value.GetSummary()
176        else:
177            answer = value.GetValue()
178        report_str = "%s expected: %s got: %s"%(expression, expected_result, answer)
179        self.assertTrue(answer == expected_result, report_str)
180
181def ApplyDecoratorsToFunction(func, decorators):
182    tmp = func
183    if type(decorators) == list:
184        for decorator in decorators:
185            tmp = decorator(tmp)
186    elif hasattr(decorators, '__call__'):
187        tmp = decorators(tmp)
188    return tmp
189
190
191def MakeInlineTest(__file, __globals, decorators=None):
192    # Adjust the filename if it ends in .pyc.  We want filenames to
193    # reflect the source python file, not the compiled variant.
194    if __file is not None and __file.endswith(".pyc"):
195        # Strip the trailing "c"
196        __file = __file[0:-1]
197
198    # Derive the test name from the current file name
199    file_basename = os.path.basename(__file)
200    InlineTest.mydir = TestBase.compute_mydir(__file)
201
202    test_name, _ = os.path.splitext(file_basename)
203    # Build the test case
204    test = type(test_name, (InlineTest,), {'using_dsym': None})
205    test.name = test_name
206
207    test.test_with_dsym = ApplyDecoratorsToFunction(test._InlineTest__test_with_dsym, decorators)
208    test.test_with_dwarf = ApplyDecoratorsToFunction(test._InlineTest__test_with_dwarf, decorators)
209    test.test_with_dwo = ApplyDecoratorsToFunction(test._InlineTest__test_with_dwo, decorators)
210
211    # Add the test case to the globals, and hide InlineTest
212    __globals.update({test_name : test})
213
214    # Keep track of the original test filename so we report it
215    # correctly in test results.
216    test.test_filename = __file
217    return test
218
219