Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
30 changes: 17 additions & 13 deletions pychunkedgraph/app/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
from .segmentation.legacy.routes import bp as segmentation_api_legacy
from .segmentation.v1.routes import bp as segmentation_api_v1
from .segmentation.generic.routes import bp as generic_api
from .app_utils import get_instance_folder_path


class CustomJsonEncoder(json.JSONEncoder):
Expand All @@ -42,10 +43,14 @@ def default(self, obj):


def create_app(test_config=None):
app = Flask(__name__)
app = Flask(
__name__,
instance_path=get_instance_folder_path(),
instance_relative_config=True,
)
app.json_encoder = CustomJsonEncoder

CORS(app, expose_headers='WWW-Authenticate')
CORS(app, expose_headers="WWW-Authenticate")

configure_app(app)

Expand All @@ -65,33 +70,32 @@ def create_app(test_config=None):

def configure_app(app):
# Load logging scheme from config.py
app_settings = os.getenv('APP_SETTINGS')
app_settings = os.getenv("APP_SETTINGS")
if not app_settings:
app.config.from_object(config.BaseConfig)
else:
app.config.from_object(app_settings)


app.config.from_pyfile("config.cfg", silent=True)
# Configure logging
# handler = logging.FileHandler(app.config['LOGGING_LOCATION'])
handler = logging.StreamHandler(sys.stdout)
handler.setLevel(app.config['LOGGING_LEVEL'])
handler.setLevel(app.config["LOGGING_LEVEL"])
formatter = jsonformatter.JsonFormatter(
fmt=app.config['LOGGING_FORMAT'],
datefmt=app.config['LOGGING_DATEFORMAT'])
fmt=app.config["LOGGING_FORMAT"], datefmt=app.config["LOGGING_DATEFORMAT"]
)
formatter.converter = time.gmtime
handler.setFormatter(formatter)
app.logger.removeHandler(default_handler)
app.logger.addHandler(handler)
app.logger.setLevel(app.config['LOGGING_LEVEL'])
app.logger.setLevel(app.config["LOGGING_LEVEL"])
app.logger.propagate = False

if app.config['USE_REDIS_JOBS']:
app.redis = redis.Redis.from_url(app.config['REDIS_URL'])
app.test_q = Queue('test', connection=app.redis)
if app.config["USE_REDIS_JOBS"]:
app.redis = redis.Redis.from_url(app.config["REDIS_URL"])
app.test_q = Queue("test", connection=app.redis)
with app.app_context():
from ..ingest.rq_cli import init_rq_cmds
from ..ingest.cli import init_ingest_cmds

init_rq_cmds(app)
init_ingest_cmds(app)
init_ingest_cmds(app)
120 changes: 113 additions & 7 deletions pychunkedgraph/app/app_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,22 +3,128 @@
from time import gmtime

import numpy as np
from flask import current_app, json
from flask import current_app, json, request
from google.auth import credentials
from google.auth import default as default_creds
from google.cloud import bigtable, datastore

from pychunkedgraph.graph import ChunkedGraph
from pychunkedgraph.logging import flask_log_db, jsonformatter
from pychunkedgraph.graph import (
exceptions as cg_exceptions,
)
from functools import wraps
from werkzeug.datastructures import ImmutableMultiDict
import time
import os

CACHE = {}


def get_app_base_path():
return os.path.dirname(os.path.realpath(__file__))


def get_instance_folder_path():
return os.path.join(get_app_base_path(), "instance")


class DoNothingCreds(credentials.Credentials):
def refresh(self, request):
pass


def remap_public(func=None, *, edit=False, check_node_ids=False):
def mydecorator(f):
@wraps(f)
def decorated_function(*args, **kwargs):
virtual_tables = current_app.config.get("VIRTUAL_TABLES", None)

# if not virtual configuration just return
if virtual_tables is None:
return f(*args, **kwargs)
table_id = kwargs.get("table_id", None)
http_args = request.args.to_dict()

if table_id is None:
# then no table remapping necessary
return f(*args, **kwargs)
if not table_id in virtual_tables:
# if table table_id isn't in virtual
# tables then just return
return f(*args, **kwargs)
else:
# then we have a virtual table
if edit:
raise cg_exceptions.Unauthorized(
"No edits allowed on virtual tables"
)
# and we want to remap the table name
new_table = virtual_tables[table_id]["table_id"]
kwargs["table_id"] = new_table
v_timestamp = virtual_tables[table_id]["timestamp"]
v_timetamp_float = time.mktime(v_timestamp.timetuple())

# we want to fix timestamp parameters too
def ceiling_timestamp(argname):
old_arg = http_args.get(argname, None)
if old_arg is not None:
# if they specified a timestamp
# enforce its less than the cap
if old_arg > v_timetamp_float:
http_args[argname] = v_timetamp_float
else:
# if they omit the timestamp, it defaults to "now"
# so we should cap it at the virtual timestamp
http_args[argname] = v_timetamp_float

ceiling_timestamp("timestamp")
ceiling_timestamp("timestamp_future")

request.args = ImmutableMultiDict(http_args)

# we also want to check for endpoints
# which ask for info about IDs and
# restrict such calls to IDs that are valid
# before the timestamp cap for this virtual table
cg = get_cg(new_table)

def assert_node_prop(prop):
node_id = kwargs.get(prop, None)
if node_id is not None:
node_id = int(node_id)
# check if this root_id is valid at this timestamp
timestamp = cg.get_node_timestamps([node_id])
if not np.all(timestamp < np.datetime64(v_timestamp)):
raise cg_exceptions.Unauthorized(
"root_id not valid at timestamp"
)

assert_node_prop("root_id")
assert_node_prop("node_id")

# some endpoints post node_ids as json, so we have to check there
# as well if the endpoint configured us to.
if check_node_ids:
node_ids = np.array(
json.loads(request.data)["node_ids"], dtype=np.uint64
)
timestamps = cg.get_node_timestamps(node_ids)
if not np.all(timestamps < np.datetime64(v_timestamp)):
raise cg_exceptions.Unauthorized(
"node_ids are all not valid at timestamp"
)

return f(*args, **kwargs)

return decorated_function

if func:
return mydecorator(func)
else:
return mydecorator


def jsonify_with_kwargs(data, as_response=True, **kwargs):
kwargs.setdefault("separators", (",", ":"))

Expand Down Expand Up @@ -113,10 +219,10 @@ def get_log_db(table_id):


def toboolean(value):
""" Transform value to boolean type.
:param value: bool/int/str
:return: bool
:raises: ValueError, if value is not boolean.
"""Transform value to boolean type.
:param value: bool/int/str
:return: bool
:raises: ValueError, if value is not boolean.
"""
if not value:
raise ValueError("Can't convert null to boolean")
Expand All @@ -137,7 +243,7 @@ def toboolean(value):


def tobinary(ids):
""" Transform id(s) to binary format
"""Transform id(s) to binary format

:param ids: uint64 or list of uint64s
:return: binary
Expand All @@ -146,7 +252,7 @@ def tobinary(ids):


def tobinary_multiples(arr):
""" Transform id(s) to binary format
"""Transform id(s) to binary format

:param arr: list of uint64 or list of uint64s
:return: binary
Expand Down
59 changes: 41 additions & 18 deletions pychunkedgraph/app/config.py
Original file line number Diff line number Diff line change
@@ -1,67 +1,90 @@
import logging
import os
import json
import datetime
from pychunkedgraph.meshing.meshgen import UTC


class BaseConfig(object):
DEBUG = False
TESTING = False
HOME = os.path.expanduser("~")
# TODO get this secret out of source control
SECRET_KEY = '1d94e52c-1c89-4515-b87a-f48cf3cb7f0b'
SECRET_KEY = "1d94e52c-1c89-4515-b87a-f48cf3cb7f0b"

LOGGING_FORMAT = '{"source":"%(name)s","time":"%(asctime)s","severity":"%(levelname)s","message":"%(message)s"}'
LOGGING_DATEFORMAT = '%Y-%m-%dT%H:%M:%S.0Z'
LOGGING_DATEFORMAT = "%Y-%m-%dT%H:%M:%S.0Z"
LOGGING_LEVEL = logging.DEBUG

CHUNKGRAPH_INSTANCE_ID = "pychunkedgraph"
PROJECT_ID = os.environ.get('PROJECT_ID', None)
CG_READ_ONLY = os.environ.get('CG_READ_ONLY', None) is not None
PCG_GRAPH_IDS = os.environ.get('PCG_GRAPH_IDS').split(",")
PROJECT_ID = os.environ.get("PROJECT_ID", None)
CG_READ_ONLY = os.environ.get("CG_READ_ONLY", None) is not None
PCG_GRAPH_IDS = os.environ.get("PCG_GRAPH_IDS").split(",")

# TODO what is this suppose to be by default?
CHUNKGRAPH_TABLE_ID = "pinky100_sv16"
# CHUNKGRAPH_TABLE_ID = "pinky100_benchmark_v92"

USE_REDIS_JOBS = False

MESHING_ENDPOINT = os.environ.get("MESHING_ENDPOINT", "http://meshing-service/meshing")
MESHING_ENDPOINT = os.environ.get(
"MESHING_ENDPOINT", "http://meshing-service/meshing"
)
daf_credential_path = os.environ.get("DAF_CREDENTIALS", None)

if daf_credential_path is not None:
with open(daf_credential_path, "r") as f:
AUTH_TOKEN = json.load(f)["token"]
else:
AUTH_TOKEN = ""


AUTH_TOKEN = ""
VIRTUAL_TABLES = {
"minnie65_public_v117": {
"table_id": "minnie3_v1",
"timestamp": datetime.datetime(
year=2021,
month=6,
day=11,
hour=8,
minute=10,
second=0,
microsecond=253,
tzinfo=datetime.timezone.utc,
),
}
}


class DevelopmentConfig(BaseConfig):
"""Development configuration."""

USE_REDIS_JOBS = False
DEBUG = True
LOGGING_LEVEL = logging.ERROR


class DockerDevelopmentConfig(DevelopmentConfig):
"""Development configuration."""

USE_REDIS_JOBS = True
REDIS_HOST = os.environ.get('REDIS_SERVICE_HOST', 'localhost')
REDIS_PORT = os.environ.get('REDIS_SERVICE_PORT', '6379')
REDIS_PASSWORD = os.environ.get('REDIS_PASSWORD', 'dev')
REDIS_URL = f'redis://:{REDIS_PASSWORD}@{REDIS_HOST}:{REDIS_PORT}/0'
REDIS_HOST = os.environ.get("REDIS_SERVICE_HOST", "localhost")
REDIS_PORT = os.environ.get("REDIS_SERVICE_PORT", "6379")
REDIS_PASSWORD = os.environ.get("REDIS_PASSWORD", "dev")
REDIS_URL = f"redis://:{REDIS_PASSWORD}@{REDIS_HOST}:{REDIS_PORT}/0"


class DeploymentWithRedisConfig(BaseConfig):
"""Deployment configuration with Redis."""

USE_REDIS_JOBS = True
REDIS_HOST = os.environ.get('REDIS_SERVICE_HOST')
REDIS_PORT = os.environ.get('REDIS_SERVICE_PORT')
REDIS_PASSWORD = os.environ.get('REDIS_PASSWORD')
REDIS_URL = f'redis://:{REDIS_PASSWORD}@{REDIS_HOST}:{REDIS_PORT}/0'
REDIS_HOST = os.environ.get("REDIS_SERVICE_HOST")
REDIS_PORT = os.environ.get("REDIS_SERVICE_PORT")
REDIS_PASSWORD = os.environ.get("REDIS_PASSWORD")
REDIS_URL = f"redis://:{REDIS_PASSWORD}@{REDIS_HOST}:{REDIS_PORT}/0"


class TestingConfig(BaseConfig):
"""Testing configuration."""

TESTING = True
USE_REDIS_JOBS = False
PRESERVE_CONTEXT_ON_EXCEPTION = False
7 changes: 6 additions & 1 deletion pychunkedgraph/app/meshing/legacy/routes.py
Original file line number Diff line number Diff line change
Expand Up @@ -3,8 +3,11 @@

from pychunkedgraph.app.meshing import common
from pychunkedgraph.graph import exceptions as cg_exceptions
from pychunkedgraph.app.app_utils import remap_public

bp = Blueprint("pcg_meshing_v0", __name__, url_prefix=f"/{common.__meshing_url_prefix__}/1.0")
bp = Blueprint(
"pcg_meshing_v0", __name__, url_prefix=f"/{common.__meshing_url_prefix__}/1.0"
)

# -------------------------------
# ------ Access control and index
Expand Down Expand Up @@ -54,6 +57,7 @@ def api_exception(e):


@bp.route("/<table_id>/<node_id>/validfragments", methods=["POST", "GET"])
@remap_public
@auth_requires_permission("view")
def handle_valid_frags(table_id, node_id):
return common.handle_valid_frags(table_id, node_id)
Expand All @@ -64,5 +68,6 @@ def handle_valid_frags(table_id, node_id):

@bp.route("/<table_id>/manifest/<node_id>:0", methods=["GET"])
@auth_requires_permission("view")
@remap_public
def handle_get_manifest(table_id, node_id):
return common.handle_get_manifest(table_id, node_id)
Loading