using ICSharpCode.SharpZipLib.Tar;
using System;
using System.Collections.Generic;
using System.ComponentModel;
using System.Diagnostics;
using System.IO;
using System.Linq;
using System.Net;
using System.Net.Http;
using System.Net.Security;
using System.Security.Authentication;
using System.Threading.Tasks;
using System.Web;
using static Tensorflow.Binding;
namespace Tensorflow.Hub
{
internal static class resolver
{
public enum ModelLoadFormat
{
[Description("COMPRESSED")]
COMPRESSED,
[Description("UNCOMPRESSED")]
UNCOMPRESSED,
[Description("AUTO")]
AUTO
}
public class DownloadManager
{
private readonly string _url;
private double _last_progress_msg_print_time;
private long _total_bytes_downloaded;
private int _max_prog_str;
private bool _interactive_mode()
{
return !string.IsNullOrEmpty(Environment.GetEnvironmentVariable("_TFHUB_DOWNLOAD_PROGRESS"));
}
private void _print_download_progress_msg(string msg, bool flush = false)
{
if (_interactive_mode())
{
// Print progress message to console overwriting previous progress
// message.
_max_prog_str = Math.Max(_max_prog_str, msg.Length);
Console.Write($"\r{msg.PadRight(_max_prog_str)}");
Console.Out.Flush();
//flushtrue
if (flush)
Console.WriteLine();
}
else
{
// Interactive progress tracking is disabled. Print progress to the
// standard TF log.
tf.Logger.Information(msg);
}
}
private void _log_progress(long bytes_downloaded)
{
// Logs progress information about ongoing module download.
_total_bytes_downloaded += bytes_downloaded;
var now = DateTime.Now.Ticks / TimeSpan.TicksPerSecond;
if (_interactive_mode() || now - _last_progress_msg_print_time > 15)
{
// Print progress message every 15 secs or if interactive progress
// tracking is enabled.
_print_download_progress_msg($"Downloading {_url}:" +
$"{tf_utils.bytes_to_readable_str(_total_bytes_downloaded, true)}");
_last_progress_msg_print_time = now;
}
}
public DownloadManager(string url)
{
_url = url;
_last_progress_msg_print_time = DateTime.Now.Ticks / TimeSpan.TicksPerSecond;
_total_bytes_downloaded = 0;
_max_prog_str = 0;
}
public void download_and_uncompress(Stream fileobj, string dst_path)
{
// Streams the content for the 'fileobj' and stores the result in dst_path.
try
{
file_utils.extract_tarfile_to_destination(fileobj, dst_path, _log_progress);
var total_size_str = tf_utils.bytes_to_readable_str(_total_bytes_downloaded, true);
_print_download_progress_msg($"Downloaded {_url}, Total size: {total_size_str}", flush: true);
}
catch (TarException ex)
{
throw new IOException($"{_url} does not appear to be a valid module. Inner message:{ex.Message}", ex);
}
}
}
private static Dictionary _flags = new();
private static readonly string _TFHUB_CACHE_DIR = "TFHUB_CACHE_DIR";
private static readonly string _TFHUB_DOWNLOAD_PROGRESS = "TFHUB_DOWNLOAD_PROGRESS";
private static readonly string _TFHUB_MODEL_LOAD_FORMAT = "TFHUB_MODEL_LOAD_FORMAT";
private static readonly string _TFHUB_DISABLE_CERT_VALIDATION = "TFHUB_DISABLE_CERT_VALIDATION";
private static readonly string _TFHUB_DISABLE_CERT_VALIDATION_VALUE = "true";
static resolver()
{
set_new_flag("tfhub_model_load_format", "AUTO");
set_new_flag("tfhub_cache_dir", null);
}
public static string model_load_format()
{
return get_env_setting(_TFHUB_MODEL_LOAD_FORMAT, "tfhub_model_load_format");
}
public static string? get_env_setting(string env_var, string flag_name)
{
string value = System.Environment.GetEnvironmentVariable(env_var);
if (string.IsNullOrEmpty(value))
{
if (_flags.ContainsKey(flag_name))
{
return _flags[flag_name];
}
else
{
return null;
}
}
else
{
return value;
}
}
public static string tfhub_cache_dir(string default_cache_dir = null, bool use_temp = false)
{
var cache_dir = get_env_setting(_TFHUB_CACHE_DIR, "tfhub_cache_dir") ?? default_cache_dir;
if (string.IsNullOrWhiteSpace(cache_dir) && use_temp)
{
// Place all TF-Hub modules under