aboutsummaryrefslogtreecommitdiffstats
path: root/contrib/python/moto/py3/moto/athena/models.py
blob: b3aef3f3c4c7968f5d3f00edb3f073fa5241234b (plain) (blame)
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
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
import time

from moto.core import BaseBackend, BaseModel, get_account_id
from moto.core.utils import BackendDict

from uuid import uuid4


class TaggableResourceMixin(object):
    # This mixing was copied from Redshift when initially implementing
    # Athena. TBD if it's worth the overhead.

    def __init__(self, region_name, resource_name, tags):
        self.region = region_name
        self.resource_name = resource_name
        self.tags = tags or []

    @property
    def arn(self):
        return "arn:aws:athena:{region}:{account_id}:{resource_name}".format(
            region=self.region,
            account_id=get_account_id(),
            resource_name=self.resource_name,
        )

    def create_tags(self, tags):
        new_keys = [tag_set["Key"] for tag_set in tags]
        self.tags = [tag_set for tag_set in self.tags if tag_set["Key"] not in new_keys]
        self.tags.extend(tags)
        return self.tags

    def delete_tags(self, tag_keys):
        self.tags = [tag_set for tag_set in self.tags if tag_set["Key"] not in tag_keys]
        return self.tags


class WorkGroup(TaggableResourceMixin, BaseModel):

    resource_type = "workgroup"
    state = "ENABLED"

    def __init__(self, athena_backend, name, configuration, description, tags):
        self.region_name = athena_backend.region_name
        super().__init__(self.region_name, "workgroup/{}".format(name), tags)
        self.athena_backend = athena_backend
        self.name = name
        self.description = description
        self.configuration = configuration


class DataCatalog(TaggableResourceMixin, BaseModel):
    def __init__(
        self, athena_backend, name, catalog_type, description, parameters, tags
    ):
        self.region_name = athena_backend.region_name
        super().__init__(self.region_name, "datacatalog/{}".format(name), tags)
        self.athena_backend = athena_backend
        self.name = name
        self.type = catalog_type
        self.description = description
        self.parameters = parameters


class Execution(BaseModel):
    def __init__(self, query, context, config, workgroup):
        self.id = str(uuid4())
        self.query = query
        self.context = context
        self.config = config
        self.workgroup = workgroup
        self.start_time = time.time()
        self.status = "QUEUED"


class NamedQuery(BaseModel):
    def __init__(self, name, description, database, query_string, workgroup):
        self.id = str(uuid4())
        self.name = name
        self.description = description
        self.database = database
        self.query_string = query_string
        self.workgroup = workgroup


class AthenaBackend(BaseBackend):
    region_name = None

    def __init__(self, region_name, account_id):
        super().__init__(region_name, account_id)
        self.work_groups = {}
        self.executions = {}
        self.named_queries = {}
        self.data_catalogs = {}

    @staticmethod
    def default_vpc_endpoint_service(service_region, zones):
        """Default VPC endpoint service."""
        return BaseBackend.default_vpc_endpoint_service_factory(
            service_region, zones, "athena"
        )

    def create_work_group(self, name, configuration, description, tags):
        if name in self.work_groups:
            return None
        work_group = WorkGroup(self, name, configuration, description, tags)
        self.work_groups[name] = work_group
        return work_group

    def list_work_groups(self):
        return [
            {
                "Name": wg.name,
                "State": wg.state,
                "Description": wg.description,
                "CreationTime": time.time(),
            }
            for wg in self.work_groups.values()
        ]

    def get_work_group(self, name):
        if name not in self.work_groups:
            return None
        wg = self.work_groups[name]
        return {
            "Name": wg.name,
            "State": wg.state,
            "Configuration": wg.configuration,
            "Description": wg.description,
            "CreationTime": time.time(),
        }

    def start_query_execution(self, query, context, config, workgroup):
        execution = Execution(
            query=query, context=context, config=config, workgroup=workgroup
        )
        self.executions[execution.id] = execution
        return execution.id

    def get_execution(self, exec_id):
        return self.executions[exec_id]

    def stop_query_execution(self, exec_id):
        execution = self.executions[exec_id]
        execution.status = "CANCELLED"

    def create_named_query(self, name, description, database, query_string, workgroup):
        nq = NamedQuery(
            name=name,
            description=description,
            database=database,
            query_string=query_string,
            workgroup=workgroup,
        )
        self.named_queries[nq.id] = nq
        return nq.id

    def get_named_query(self, query_id):
        return self.named_queries[query_id] if query_id in self.named_queries else None

    def list_data_catalogs(self):
        return [
            {"CatalogName": dc.name, "Type": dc.type}
            for dc in self.data_catalogs.values()
        ]

    def get_data_catalog(self, name):
        if name not in self.data_catalogs:
            return None
        dc = self.data_catalogs[name]
        return {
            "Name": dc.name,
            "Description": dc.description,
            "Type": dc.type,
            "Parameters": dc.parameters,
        }

    def create_data_catalog(self, name, catalog_type, description, parameters, tags):
        if name in self.data_catalogs:
            return None
        data_catalog = DataCatalog(
            self, name, catalog_type, description, parameters, tags
        )
        self.data_catalogs[name] = data_catalog
        return data_catalog


athena_backends = BackendDict(AthenaBackend, "athena")